From a838911a55840943d974ad35330d7717cc14f14d Mon Sep 17 00:00:00 2001 From: younesbelkada Date: Mon, 9 Dec 2024 16:45:31 +0000 Subject: [PATCH] more changes to support chat models --- run_inference.py | 4 +++- setup_env.py | 3 +++ utils/convert-hf-to-gguf-bitnet.py | 2 +- 3 files changed, 7 insertions(+), 2 deletions(-) diff --git a/run_inference.py b/run_inference.py index fe14e0e..7c279cd 100644 --- a/run_inference.py +++ b/run_inference.py @@ -30,7 +30,8 @@ def run_inference(): '-ngl', '0', '-c', str(args.ctx_size), '--temp', str(args.temperature), - "-b", "1" + "-b", "1", + "-cnv" if args.cnv else "" ] run_command(command) @@ -48,6 +49,7 @@ if __name__ == "__main__": parser.add_argument("-t", "--threads", type=int, help="Number of threads to use", required=False, default=2) parser.add_argument("-c", "--ctx-size", type=int, help="Size of the prompt context", required=False, default=2048) parser.add_argument("-temp", "--temperature", type=float, help="Temperature, a hyperparameter that controls the randomness of the generated text", required=False, default=0.8) + parser.add_argument("-cnv", "--conversation", ction='store_true', help="Whether to enable chat mode or not (for instruct models.)") args = parser.parse_args() run_inference() \ No newline at end of file diff --git a/setup_env.py b/setup_env.py index ab4beb9..1ea456b 100644 --- a/setup_env.py +++ b/setup_env.py @@ -20,6 +20,9 @@ SUPPORTED_HF_MODELS = { "HF1BitLLM/Llama3-8B-1.58-100B-tokens": { "model_name": "Llama3-8B-1.58-100B-tokens", }, + "tiiuae/falcon3-7b-instruct-1.58bit": { + "model_name": "falcon3-7b-1.58bit", + }, "tiiuae/falcon3-7b-1.58bit": { "model_name": "falcon3-7b-1.58bit", } diff --git a/utils/convert-hf-to-gguf-bitnet.py b/utils/convert-hf-to-gguf-bitnet.py index 8faa51e..f525f58 100644 --- a/utils/convert-hf-to-gguf-bitnet.py +++ b/utils/convert-hf-to-gguf-bitnet.py @@ -676,7 +676,7 @@ def read_model_config(model_dir: str) -> dict[str, Any]: with open(config, "r") as f: return json.load(f) -@Model.register("LLaMAForCausalLM", "LlamaForCausalLM", "MistralForCausalLM", "MixtralForCausalLM", "Falcon3ForCausalLM") +@Model.register("LLaMAForCausalLM", "LlamaForCausalLM", "MistralForCausalLM", "MixtralForCausalLM") class LlamaModel(Model): model_arch = gguf.MODEL_ARCH.LLAMA