diff --git a/koboldcpp.py b/koboldcpp.py index 73b893edd..900f5bd71 100644 --- a/koboldcpp.py +++ b/koboldcpp.py @@ -110,6 +110,7 @@ importvars_in_progress = False has_multiplayer = False has_audio_support = False has_vision_support = False +cached_chat_template = None savedata_obj = None multiplayer_story_data_compressed = None #stores the full compressed story of the current multiplayer session multiplayer_turn_major = 1 # to keep track of when a client needs to sync their stories @@ -2332,6 +2333,22 @@ def is_ipv6_supported(): except Exception: return False +def format_jinja(messages): + try: + def strftime_now(format='%Y-%m-%d %H:%M:%S'): + return datetime.now().strftime(format) + global cached_chat_template + from jinja2.sandbox import ImmutableSandboxedEnvironment + jinja_env = ImmutableSandboxedEnvironment(trim_blocks=True, lstrip_blocks=True) + jinja_env.globals['strftime_now'] = strftime_now + jinja_compiled_template = jinja_env.from_string(cached_chat_template) + text = jinja_compiled_template.render(messages=messages, add_generation_prompt=True, bos_token="", eos_token="") + return text if text else None + except Exception as e: + print("Jinja formatting failed: {e}") + return None + + # Used to parse json for openai tool calls def extract_json_from_string(input_string): parsed_json = None @@ -2504,7 +2521,7 @@ def determine_tool_json_to_use(genparams, curr_ctx, assistant_message_start, is_ return used_tool_json -def transform_genparams(genparams, api_format): +def transform_genparams(genparams, api_format, use_jinja): global chatcompl_adapter, maxctx if api_format < 0: #not text gen, do nothing @@ -2643,85 +2660,91 @@ ws ::= | " " | "\n" [ \t]{0,20} message_index = 0 attachedimgid = 0 attachedaudid = 0 - for message in messages_array: - message_index += 1 - if message['role'] == "system": - messages_string += system_message_start - elif message['role'] == "user": - messages_string += user_message_start - elif message['role'] == "assistant": - messages_string += assistant_message_start - elif message['role'] == "tool": - messages_string += tools_message_start + jinja_output = None + if use_jinja and cached_chat_template: + jinja_output = format_jinja(messages_array) + if jinja_output: + messages_string = jinja_output + else: + for message in messages_array: + message_index += 1 + if message['role'] == "system": + messages_string += system_message_start + elif message['role'] == "user": + messages_string += user_message_start + elif message['role'] == "assistant": + messages_string += assistant_message_start + elif message['role'] == "tool": + messages_string += tools_message_start - # content can be a string or an array of objects - curr_content = message.get("content",None) - if api_format==7: #ollama handle vision - imgs = message.get("images",None) - if imgs and len(imgs) > 0: - for img in imgs: - images_added.append(img) - if not curr_content: - if "tool_calls" in message: - try: - if len(message.get("tool_calls"))>0: - tcfnname = message.get("tool_calls")[0].get("function").get("name") - messages_string += f"\n(Made a function call to {tcfnname})\n" - except Exception: - messages_string += "\n(Made a function call)\n" - pass # do nothing - elif isinstance(curr_content, str): - messages_string += curr_content - elif isinstance(curr_content, list): #is an array - for item in curr_content: - if item['type']=="text": - messages_string += item['text'] - elif item['type']=="image_url": - if 'image_url' in item and item['image_url'] and item['image_url']['url'] and item['image_url']['url'].startswith("data:image"): - images_added.append(item['image_url']['url'].split(",", 1)[1]) - attachedimgid += 1 - messages_string += f"\n(Attached Image {attachedimgid})\n" - elif item['type']=="input_audio": - if 'input_audio' in item and item['input_audio'] and item['input_audio']['data']: - audio_added.append(item['input_audio']['data']) - attachedaudid += 1 - messages_string += f"\n(Attached Audio {attachedaudid})\n" - # If last message, add any tools calls after message content and before message end token if any - if message_index == len(messages_array): - used_tool_json = determine_tool_json_to_use(genparams, messages_string, assistant_message_start, (message['role'] == "tool")) + # content can be a string or an array of objects + curr_content = message.get("content",None) + if api_format==7: #ollama handle vision + imgs = message.get("images",None) + if imgs and len(imgs) > 0: + for img in imgs: + images_added.append(img) + if not curr_content: + if "tool_calls" in message: + try: + if len(message.get("tool_calls"))>0: + tcfnname = message.get("tool_calls")[0].get("function").get("name") + messages_string += f"\n(Made a function call to {tcfnname})\n" + except Exception: + messages_string += "\n(Made a function call)\n" + pass # do nothing + elif isinstance(curr_content, str): + messages_string += curr_content + elif isinstance(curr_content, list): #is an array + for item in curr_content: + if item['type']=="text": + messages_string += item['text'] + elif item['type']=="image_url": + if 'image_url' in item and item['image_url'] and item['image_url']['url'] and item['image_url']['url'].startswith("data:image"): + images_added.append(item['image_url']['url'].split(",", 1)[1]) + attachedimgid += 1 + messages_string += f"\n(Attached Image {attachedimgid})\n" + elif item['type']=="input_audio": + if 'input_audio' in item and item['input_audio'] and item['input_audio']['data']: + audio_added.append(item['input_audio']['data']) + attachedaudid += 1 + messages_string += f"\n(Attached Audio {attachedaudid})\n" + # If last message, add any tools calls after message content and before message end token if any + if message_index == len(messages_array): + used_tool_json = determine_tool_json_to_use(genparams, messages_string, assistant_message_start, (message['role'] == "tool")) - if used_tool_json: - toolparamjson = None - toolname = None - # Set temperature lower automatically if function calling, cannot exceed 0.5 - genparams["temperature"] = (1.0 if genparams.get("temperature", 0.5) > 1.0 else genparams.get("temperature", 0.5)) - genparams["using_openai_tools"] = True - # Set grammar to llamacpp example grammar to force json response (see https://github.com/ggerganov/llama.cpp/blob/master/grammars/json_arr.gbnf) - genparams["grammar"] = jsongrammar - try: - toolname = used_tool_json.get('function').get('name') - toolparamjson = used_tool_json.get('function').get('parameters') - bettergrammarjson = {"type":"array","items":{"type":"object","properties":{"id":{"type":"string","enum":["call_001"]},"type":{"type":"string","enum":["function"]},"function":{"type":"object","properties":{"name":{"type":"string"},"arguments":{}},"required":["name","arguments"],"additionalProperties":False}},"required":["id","type","function"],"additionalProperties":False}} - bettergrammarjson["items"]["properties"]["function"]["properties"]["arguments"] = toolparamjson - decoded = convert_json_to_gbnf(bettergrammarjson) - if decoded: - genparams["grammar"] = decoded - except Exception: - pass - tool_json_formatting_instruction = f"\nPlease use the provided schema to fill the parameters to create a function call for {toolname}, in the following format: " + json.dumps([{"id": "call_001", "type": "function", "function": {"name": f"{toolname}", "arguments": {"first property key": "first property value", "second property key": "second property value"}}}], indent=0) - messages_string += f"\n\nJSON Schema:\n{used_tool_json}\n\n{tool_json_formatting_instruction}{assistant_message_start}" + if used_tool_json: + toolparamjson = None + toolname = None + # Set temperature lower automatically if function calling, cannot exceed 0.5 + genparams["temperature"] = (1.0 if genparams.get("temperature", 0.5) > 1.0 else genparams.get("temperature", 0.5)) + genparams["using_openai_tools"] = True + # Set grammar to llamacpp example grammar to force json response (see https://github.com/ggerganov/llama.cpp/blob/master/grammars/json_arr.gbnf) + genparams["grammar"] = jsongrammar + try: + toolname = used_tool_json.get('function').get('name') + toolparamjson = used_tool_json.get('function').get('parameters') + bettergrammarjson = {"type":"array","items":{"type":"object","properties":{"id":{"type":"string","enum":["call_001"]},"type":{"type":"string","enum":["function"]},"function":{"type":"object","properties":{"name":{"type":"string"},"arguments":{}},"required":["name","arguments"],"additionalProperties":False}},"required":["id","type","function"],"additionalProperties":False}} + bettergrammarjson["items"]["properties"]["function"]["properties"]["arguments"] = toolparamjson + decoded = convert_json_to_gbnf(bettergrammarjson) + if decoded: + genparams["grammar"] = decoded + except Exception: + pass + tool_json_formatting_instruction = f"\nPlease use the provided schema to fill the parameters to create a function call for {toolname}, in the following format: " + json.dumps([{"id": "call_001", "type": "function", "function": {"name": f"{toolname}", "arguments": {"first property key": "first property value", "second property key": "second property value"}}}], indent=0) + messages_string += f"\n\nJSON Schema:\n{used_tool_json}\n\n{tool_json_formatting_instruction}{assistant_message_start}" - if message['role'] == "system": - messages_string += system_message_end - elif message['role'] == "user": - messages_string += user_message_end - elif message['role'] == "assistant": - messages_string += assistant_message_end - elif message['role'] == "tool": - messages_string += tools_message_end - - messages_string += assistant_message_gen + if message['role'] == "system": + messages_string += system_message_end + elif message['role'] == "user": + messages_string += user_message_end + elif message['role'] == "assistant": + messages_string += assistant_message_end + elif message['role'] == "tool": + messages_string += tools_message_end + messages_string += assistant_message_gen + genparams["prompt"] = messages_string if len(images_added)>0: genparams["images"] = images_added @@ -3370,7 +3393,7 @@ Change Mode
def do_GET(self): global embedded_kailite, embedded_kcpp_docs, embedded_kcpp_sdui, embedded_kailite_gz, embedded_kcpp_docs_gz, embedded_kcpp_sdui_gz, embedded_lcpp_ui_gz - global last_req_time, start_time + global last_req_time, start_time, cached_chat_template global savedata_obj, has_multiplayer, multiplayer_turn_major, multiplayer_turn_minor, multiplayer_story_data_compressed, multiplayer_dataformat, multiplayer_lastactive, maxctx, maxhordelen, friendlymodelname, lastuploadedcomfyimg, lastgeneratedcomfyimg, KcppVersion, totalgens, preloaded_story, exitcounter, currentusergenkey, friendlysdmodelname, fullsdmodelpath, password, friendlyembeddingsmodelname self.path = self.path.rstrip('/') response_body = None @@ -3596,11 +3619,9 @@ Change Mode
elif self.path.endswith(('/.well-known/serviceinfo')): response_body = (json.dumps({"version":"0.2","software":{"name":"KoboldCpp","version":KcppVersion,"repository":"https://github.com/LostRuins/koboldcpp","homepage":"https://github.com/LostRuins/koboldcpp","logo":"https://raw.githubusercontent.com/LostRuins/koboldcpp/refs/heads/concedo/niko.ico"},"api":{"koboldai":{"name":"KoboldAI API","rel_url":"/api","documentation":"https://lite.koboldai.net/koboldcpp_api","version":KcppVersion},"openai":{"name":"OpenAI API","rel_url ":"/v1","documentation":"https://openai.com/documentation/api","version":KcppVersion}}}).encode()) - elif self.path=="/props": - ctbytes = handle.get_chat_template() - chat_template = ctypes.string_at(ctbytes).decode("UTF-8","ignore") + elif self.path=="/props": response_body = (json.dumps({ - "chat_template": chat_template, + "chat_template": cached_chat_template, "id": 0, "id_task": -1, "total_slots": 1, @@ -4101,6 +4122,7 @@ Change Mode
is_tts = False is_embeddings = False response_body = None + use_jinja = args.jinja if self.path.endswith('/api/admin/check_state'): if global_memory and args.admin and args.admindir and os.path.exists(args.admindir) and self.check_header_password(args.adminpassword): @@ -4240,7 +4262,7 @@ Change Mode
utfprint("\nInput: " + json.dumps(printablegenparams_raw,ensure_ascii=False),1) # transform genparams (only used for text gen) first - genparams = transform_genparams(genparams, api_format) + genparams = transform_genparams(genparams, api_format, use_jinja) if args.debugmode >= 1: printablegenparams = truncate_long_json(genparams,trunc_len) @@ -4820,6 +4842,7 @@ def show_gui(): customrope_base = ctk.StringVar(value="10000") customrope_nativectx = ctk.StringVar(value=str(default_native_ctx)) chatcompletionsadapter_var = ctk.StringVar(value="AutoGuess") + jinja_var = ctk.IntVar(value=0) moeexperts_var = ctk.StringVar(value=str(-1)) moecpu_var = ctk.StringVar(value=str(0)) defaultgenamt_var = ctk.StringVar(value=str(768)) @@ -5532,11 +5555,12 @@ def show_gui(): quantkv_var.trace_add("write", toggleflashattn) makecheckbox(tokens_tab, "No BOS Token", nobostoken_var, 43, tooltiptxt="Prevents BOS token from being added at the start of any prompt. Usually NOT recommended for most models.") makecheckbox(tokens_tab, "Enable Guidance", enableguidance_var, 43,padx=(200 if corrupt_scaler else 140), tooltiptxt="Enables the use of Classifier-Free-Guidance, which allows the use of negative prompts. Has performance and memory impact.") + makecheckbox(tokens_tab, "Use Jinja", jinja_var, row=45, tooltiptxt="Enables using jinja chat template formatting for chat completions endpoint. Other endpoints are unaffected.") makelabelentry(tokens_tab, "MoE Experts:", moeexperts_var, row=55, padx=(220 if corrupt_scaler else 120), singleline=True, tooltip="Override number of MoE experts.") makelabelentry(tokens_tab, "MoE CPU Layers:", moecpu_var, row=55, padx=(490 if corrupt_scaler else 320), singleline=True, tooltip="Force Mixture of Experts (MoE) weights of the first N layers to the CPU.\nSetting it higher than GPU layers has no effect.", labelpadx=(300 if corrupt_scaler else 210)) makelabelentry(tokens_tab, "Override KV:", override_kv_var, row=57, padx=(220 if corrupt_scaler else 120), singleline=True, width=150, tooltip="Override metadata value by key. Separate multiple values with commas. Format is name=type:value. Types: int, float, bool, str") makelabelentry(tokens_tab, "Override Tensors:", override_tensors_var, row=59, padx=(220 if corrupt_scaler else 120), singleline=True, width=150, tooltip="Override selected backend for specific tensors matching tensor_name_regex_pattern=buffer_type, same as in llama.cpp.") - + # Model Tab model_tab = tabcontent["Loaded Files"] @@ -5843,6 +5867,7 @@ def show_gui(): args.defaultgenamt = int(defaultgenamt_var.get()) if defaultgenamt_var.get()!="" else 768 args.genlimit = int(genlimit_var.get()) if genlimit_var.get()!="" else 0 args.nobostoken = (nobostoken_var.get()==1) + args.jinja = (jinja_var.get()==1) args.enableguidance = (enableguidance_var.get()==1) args.overridekv = None if override_kv_var.get() == "" else override_kv_var.get() args.overridetensors = None if override_tensors_var.get() == "" else override_tensors_var.get() @@ -6082,6 +6107,7 @@ def show_gui(): else: genlimit_var.set(str(0)) nobostoken_var.set(dict["nobostoken"] if ("nobostoken" in dict) else 0) + jinja_var.set(dict["jinja"] if ("jinja" in dict) else 0) enableguidance_var.set(dict["enableguidance"] if ("enableguidance" in dict) else 0) if "overridekv" in dict and dict["overridekv"]: override_kv_var.set(dict["overridekv"]) @@ -7132,7 +7158,7 @@ def main(launch_args, default_args): def kcpp_main_process(launch_args, g_memory=None, gui_launcher=False): global embedded_kailite, embedded_kcpp_docs, embedded_kcpp_sdui, embedded_kailite_gz, embedded_kcpp_docs_gz, embedded_kcpp_sdui_gz, embedded_lcpp_ui_gz, start_time, exitcounter, global_memory, using_gui_launcher - global libname, args, friendlymodelname, friendlysdmodelname, fullsdmodelpath, password, fullwhispermodelpath, ttsmodelpath, embeddingsmodelpath, friendlyembeddingsmodelname, has_audio_support, has_vision_support + global libname, args, friendlymodelname, friendlysdmodelname, fullsdmodelpath, password, fullwhispermodelpath, ttsmodelpath, embeddingsmodelpath, friendlyembeddingsmodelname, has_audio_support, has_vision_support, cached_chat_template start_server = True @@ -7502,16 +7528,15 @@ def kcpp_main_process(launch_args, g_memory=None, gui_launcher=False): if not loadok: exitcounter = 999 exit_with_error(3,"Could not load text model: " + modelname) - - if (chatcompl_adapter_list is not None and isinstance(chatcompl_adapter_list, list)): + # The chat completions adapter is a list that needs derivation from chat templates # Try to derive chat completions adapter from chat template, now that we have the model loaded if not args.nomodel and args.model_param: ctbytes = handle.get_chat_template() - chat_template = ctypes.string_at(ctbytes).decode("UTF-8","ignore") - if chat_template != "": + cached_chat_template = ctypes.string_at(ctbytes).decode("UTF-8","ignore") + if cached_chat_template != "" and (chatcompl_adapter_list is not None and isinstance(chatcompl_adapter_list, list)): for entry in chatcompl_adapter_list: - if all(s in chat_template for s in entry['search']): + if all(s in cached_chat_template for s in entry['search']): print(f"Chat completion heuristic: {entry['name']}") chatcompl_adapter = entry['adapter'] break @@ -7785,7 +7810,7 @@ def kcpp_main_process(launch_args, g_memory=None, gui_launcher=False): continue lastturns.append({"role":"user","content":lastuserinput}) payload = {"messages":lastturns,"rep_pen":1.07,"temperature":0.8} - payload = transform_genparams(payload, 4) #to chat completions + payload = transform_genparams(payload, 4, False) #to chat completions if args.debugmode < 1: suppress_stdout() genout = generate(genparams=payload) @@ -7969,6 +7994,7 @@ if __name__ == '__main__': advparser.add_argument("--ratelimit", metavar=('[seconds]'), help="If enabled, rate limit generative request by IP address. Each IP can only send a new request once per X seconds.", type=int, default=0) advparser.add_argument("--ignoremissing", help="Ignores all missing non-essential files, just skipping them instead.", action='store_true') advparser.add_argument("--chatcompletionsadapter", metavar=('[filename]'), help="Select an optional ChatCompletions Adapter JSON file to force custom instruct tags.", default="AutoGuess") + advparser.add_argument("--jinja", help="Enables using jinja chat template formatting for chat completions endpoint. Other endpoints are unaffected.", action='store_true') advparser.add_argument("--flashattention","--flash-attn","-fa", help="Enables flash attention.", action='store_true') advparser.add_argument("--lowvram","-nkvo","--no-kv-offload", help="If supported by the backend, do not offload KV to GPU (lowvram mode). Not recommended, will be slow.", action='store_true') advparser.add_argument("--quantkv", help="Sets the KV cache data type quantization, 0=f16, 1=q8, 2=q4. Requires Flash Attention for full effect, otherwise only K cache is quantized.",metavar=('[quantization level 0/1/2]'), type=int, choices=[0,1,2], default=0) diff --git a/requirements.txt b/requirements.txt index 86c701ce8..45207f264 100644 --- a/requirements.txt +++ b/requirements.txt @@ -5,4 +5,5 @@ gguf>=0.1.0 customtkinter>=5.1.0 protobuf>=4.21.0 psutil>=5.9.4 -darkdetect>=0.8.0 \ No newline at end of file +darkdetect>=0.8.0 +jinja2>=3.1 \ No newline at end of file