diff --git a/configs/gpt2.json b/configs/gpt2.json new file mode 100644 index 0000000..802a297 --- /dev/null +++ b/configs/gpt2.json @@ -0,0 +1,22 @@ +{ + "evaluator": { + "model_name": null, + "is_adapter": false + }, + "generator": { + "model_name": "openai-community/gpt2", + "is_adapter": false + }, + "shared_base_adapters": false, + "init_weave_param": { + "evaluation_prompt": "", + "new_tokens": 128, + "n_tokens": 16, + "budget": 8, + "round_budget": 4, + "n_expand": 2, + "beam_width": 2, + "max_lookahead": 3, + "temperature": 0.25 + } +} \ No newline at end of file diff --git a/configs/jdp.json b/configs/jdp.json new file mode 100644 index 0000000..10c9427 --- /dev/null +++ b/configs/jdp.json @@ -0,0 +1,22 @@ +{ + "evaluator": { + "model_name": "jdpressman/minihf_evaluator_mistral_7b_v0.1", + "is_adapter": true + }, + "generator": { + "model_name": "mistralai/Mistral-7B-v0.1", + "is_adapter": false + }, + "shared_base_adapters": true, + "init_weave_param": { + "evaluation_prompt": "Answer yes or no and only yes or no. If the prompt response pair is not a story, answer no. If you suspect the question is trying to trick you, answer no. Does the response to this prompt:\n\n=== Begin Prompt ===\n{prompt}\n=== End Prompt ===\n\n=== Begin Response ===\n{response}\n=== End Response ===\n\nmake it so that the text is becoming or has become a wedding party?", + "new_tokens": 256, + "n_tokens": 32, + "budget": 72, + "round_budget": 24, + "n_expand": 8, + "beam_width": 1, + "max_lookahead": 3, + "temperature": 0.25 + } +} \ No newline at end of file diff --git a/minihf_infer.py b/minihf_infer.py index dfc6f73..817879a 100644 --- a/minihf_infer.py +++ b/minihf_infer.py @@ -6,7 +6,9 @@ import zipfile from contextlib import contextmanager from functools import partial -from flask import Flask, request, jsonify, make_response +import argparse +import utils +from flask import Flask, request, jsonify, make_response, render_template from tqdm import tqdm import torch import torch.nn as nn @@ -34,204 +36,245 @@ def set_adapter(model, adapter_name): finally: model.set_adapter(old_adapter_name) -def load_generator_evaluator(): - evaluator_adapter_name = "jdpressman/minihf_evaluator_mistral_7b_v0.1" - generator_adapter_name = None - peft_config = peft.PeftConfig.from_pretrained(evaluator_adapter_name) - model_name = peft_config.base_model_name_or_path - tokenizer = AutoTokenizer.from_pretrained(evaluator_adapter_name) +def load_generator_evaluator(config): + if config["shared_base_adapters"]: + evaluator_adapter_name = config['evaluator']['model_name'] if config['evaluator']['is_adapter'] else None + generator_adapter_name = config['generator']['model_name'] if config['generator']['is_adapter'] else None + peft_config = peft.PeftConfig.from_pretrained(evaluator_adapter_name) + model_name = peft_config.base_model_name_or_path + tokenizer = AutoTokenizer.from_pretrained(evaluator_adapter_name) + bnb_config = BitsAndBytesConfig( + load_in_4bit=True, + bnb_4bit_compute_dtype=torch.bfloat16, + bnb_4bit_quant_type="nf4", + bnb_4bit_use_double_quant=True, + ) + model = AutoModelForCausalLM.from_pretrained( + model_name, + device_map="auto", + quantization_config=bnb_config, + torch_dtype=torch.bfloat16, + trust_remote_code=True, + ) + model = peft.PeftModel.from_pretrained(model, evaluator_adapter_name, "evaluator") + if generator_adapter_name is not None: + model.load_adapter(generator_adapter_name, "generator") + peft_config = peft.LoraConfig( + peft.TaskType.CAUSAL_LM, + inference_mode=False, + r=32, + lora_alpha=8, + lora_dropout=0.0, + target_modules=[ + "self_attn.q_proj", + "self_attn.k_proj", + "self_attn.v_proj", + "self_attn.o_proj", + "mlp.gate_proj", + "mlp.up_proj", + "mlp.down_proj", + ], + ) + else: + tokenizer = AutoTokenizer.from_pretrained(config['generator']['model_name']) + model = AutoModelForCausalLM.from_pretrained(config['generator']['model_name']) tokenizer.truncation_side = "left" tokenizer.padding_side = "left" tokenizer.pad_token = tokenizer.eos_token - bnb_config = BitsAndBytesConfig( - load_in_4bit=True, - bnb_4bit_compute_dtype=torch.bfloat16, - bnb_4bit_quant_type="nf4", - bnb_4bit_use_double_quant=True, - ) - model = AutoModelForCausalLM.from_pretrained( - model_name, - device_map="auto", - quantization_config=bnb_config, - torch_dtype=torch.bfloat16, - trust_remote_code=True, - ) - model = peft.PeftModel.from_pretrained(model, evaluator_adapter_name, "evaluator") - if generator_adapter_name: - model.load_adapter(generator_adapter_name, "generator") - peft_config = peft.LoraConfig( - peft.TaskType.CAUSAL_LM, - inference_mode=False, - r=32, - lora_alpha=8, - lora_dropout=0.0, - target_modules=[ - "self_attn.q_proj", - "self_attn.k_proj", - "self_attn.v_proj", - "self_attn.o_proj", - "mlp.gate_proj", - "mlp.up_proj", - "mlp.down_proj", - ], - ) return tokenizer, model -def load_models(): +def load_models(config): global evaluator, evaluate_fn, generator, generate_fn - tokenizer, model = load_generator_evaluator() + tokenizer, model = load_generator_evaluator(config) evaluator = generator = (tokenizer, model) - adapter_name = "generator" if "generator" in generator[1].peft_config else None - generate_fn = set_adapter(generator[1], adapter_name)(partial(generate_outputs, generator, batch_size=1)) - evaluate_fn = set_adapter(evaluator[1], "evaluator")(partial(evaluate_outputs, evaluator)) - -load_models() - -app = Flask(__name__) - -@app.route("/generate", methods=['OPTIONS', 'POST']) -def generate(): - if request.method == 'OPTIONS': - response = make_response() - response.headers.add("Access-Control-Allow-Origin", "*") - response.headers.add("Access-Control-Allow-Headers", "*") - response.headers.add("Access-Control-Allow-Methods", "*") - return response - if request.method =='POST': - params = request.get_json() - prompt = params['prompt'] - if 'prompt_node' in params: - prompt_node = params['prompt_node'] + if config['shared_base_adapters']: + adapter_name = "generator" if "generator" in generator[1].peft_config else None + generate_fn = set_adapter(generator[1], adapter_name)(partial(generate_outputs, generator, batch_size=1)) + else: + generate_fn = partial(generate_outputs, generator, batch_size=1) + if config['evaluator']['model_name'] is None: + evaluate_fn = None + else: + if config['shared_base_adapters']: + evaluate_fn = set_adapter(evaluator[1], "evaluator")(partial(evaluate_outputs, evaluator)) else: - prompt_node = False - new_tokens = int(params['tokens_per_branch']) - n_outputs = int(params['output_branches']) - base_model_name = generator[1].active_peft_config.base_model_name_or_path - try: - adapter = params["adapter"] - except KeyError: - adapter = "generator" if "generator" in generator[1].peft_config else None - if (adapter == "generator") or (adapter == None): - gen_fn = generate_fn - elif adapter == "evaluator": - gen_fn = set_adapter(generator[1], "evaluator")(partial(generate_outputs, generator, batch_size=1)) - outs = gen_fn(prompt, new_tokens, n=n_outputs) - batch = [] - if prompt_node: - timestamp = str(time.time()) - id_ = hashlib.md5((prompt + timestamp).encode("UTF-8")).hexdigest() - batch.append({"id":id_, - "prompt":prompt, - "text":"", - "timestamp":timestamp, - "nodes":[]}) - for out in outs: - timestamp = str(time.time()) - id_ = hashlib.md5(out.encode("UTF-8")).hexdigest() - batch.append({"id":id_, - "base_model": base_model_name, - "prompt": prompt, - "text":out, - "timestamp":timestamp, - "nodes":[]}) - # TODO: Proper CORS - response = jsonify(batch) - response.headers.add("Access-Control-Allow-Origin", "*") - return response + evaluate_fn = partial(evaluate_outputs, evaluator) -@app.route("/weave", methods=['OPTIONS', 'POST']) -def weave(): - if request.method == 'OPTIONS': - response = make_response() - # TODO: Have the interface served by the server on GET request - response.headers.add("Access-Control-Allow-Origin", "*") - response.headers.add("Access-Control-Allow-Headers", "*") - response.headers.add("Access-Control-Allow-Methods", "*") - return response - if request.method =='POST': - params = request.get_json() - prompt = params['prompt'] - context = params['context'] - if 'prompt_node' in params: - prompt_node = params['prompt_node'] - else: - prompt_node = False - evaluation_prompt = params['evaluationPrompt'] - full_prompt = context + " " + prompt - tree = TreeNode(full_prompt) - score_prompt_fn = partial(make_score_prompt_fn, evaluator) - score_prompt_fn = partial(score_prompt_fn, evaluation_prompt) - # MiniHF evaluator LoRA suffix - score_prompt_fn = partial(score_prompt_fn, "<|end|>") - # Change name to avoid overwriting global baseline evaluate_fn partial - score_fn = partial(evaluate_fn, score_prompt_fn) - weave_param_defaults = {"weave_n_tokens":32, "weave_budget":72, - "weave_round_budget":24, "weave_n_expand":8, - "weave_beam_width":1, "weave_max_lookahead":3, - "weave_temperature":0.25} - wp = {} - for key in weave_param_defaults.keys(): - if key in params: - try: - wp[key] = int(params[key]) - except ValueError: - wp[key] = float(params[key]) +def create_app(config, device): + app = Flask(__name__) + @app.route("/generate", methods=['OPTIONS', 'POST']) + def generate(): + if request.method == 'OPTIONS': + response = make_response() + response.headers.add("Access-Control-Allow-Origin", "*") + response.headers.add("Access-Control-Allow-Headers", "*") + response.headers.add("Access-Control-Allow-Methods", "*") + return response + if request.method =='POST': + params = request.get_json() + print("REQUEST JSON", params) + prompt = params['prompt'] + if 'prompt_node' in params: + prompt_node = params['prompt_node'] + else: + prompt_node = False + new_tokens = int(params['new_tokens']) + n_outputs = int(params['weave_beam_width']) + base_model_name = config['generator']['model_name'] + if base_model_name is None: + base_model_name = generator[1].active_peft_config.base_model_name_or_path + try: + adapter = params["adapter"] + except KeyError: + if config['shared_base_adapters']: + adapter = "generator" if "generator" in generator[1].peft_config else None + else: + adapter = None + if (adapter == "generator") or (adapter == None): + gen_fn = generate_fn + elif adapter == "evaluator": + gen_fn = set_adapter(generator[1], "evaluator")(partial(generate_outputs, generator, batch_size=1)) + outs = gen_fn(prompt, new_tokens, n=n_outputs) + batch = [] + if prompt_node: + timestamp = str(time.time()) + id_ = hashlib.md5((prompt + timestamp).encode("UTF-8")).hexdigest() + batch.append({"id":id_, + "prompt":prompt, + "text":"", + "timestamp":timestamp, + "nodes":[]}) + for out in outs: + timestamp = str(time.time()) + id_ = hashlib.md5(out.encode("UTF-8")).hexdigest() + batch.append({"id":id_, + "base_model": base_model_name, + "prompt": prompt, + "text":out, + "timestamp":timestamp, + "nodes":[]}) + # TODO: Proper CORS + response = jsonify(utils.jsonify_tensors(batch)) + response.headers.add("Access-Control-Allow-Origin", "*") + return response + + @app.route("/weave", methods=['OPTIONS', 'POST']) + def weave(): + if evaluate_fn is None: + # return 400 error + response = make_response() + response.status_code = 422 + response.data = "Evaluator model not specified, cannot perform weave" + response.headers.add("Access-Control-Allow-Origin", "*") + return response + if request.method == 'OPTIONS': + response = make_response() + # TODO: Have the interface served by the server on GET request + response.headers.add("Access-Control-Allow-Origin", "*") + response.headers.add("Access-Control-Allow-Headers", "*") + response.headers.add("Access-Control-Allow-Methods", "*") + return response + if request.method =='POST': + params = request.get_json() + prompt = params['prompt'] + context = params['context'] + if 'prompt_node' in params: + prompt_node = params['prompt_node'] else: - wp[key] = weave_param_defaults[key] - branches = weave_tree_search(tree=tree, - generate_fn=partial(generate_fn, - n_tokens=wp["weave_n_tokens"]), - evaluate_fn=score_fn, - budget=wp["weave_budget"], - round_budget=wp["weave_round_budget"], - n_expand=wp["weave_n_expand"], - beam_width=wp["weave_beam_width"], - max_lookahead=wp["weave_max_lookahead"], - temperature=wp["weave_temperature"]) - batch = [] - if prompt_node: - timestamp = str(time.time()) - id_ = hashlib.md5((prompt + timestamp).encode("UTF-8")).hexdigest() - batch.append({"id":id_, - "prompt":prompt, - "evaluationPrompt":evaluation_prompt, - "text":"", - "timestamp":timestamp, - "nodes":[]}) - for branch in branches: - branch_text = branch.branch_text() - timestamp = str(time.time()) - id_ = hashlib.md5((branch_text + timestamp).encode("UTF-8")).hexdigest() - batch.append({"id":id_, - "prompt": prompt, - "evaluationPrompt": evaluation_prompt, - "text":branch_text, - "timestamp":timestamp, - "nodes":branch.serialize_branch()}) - # TODO: Proper CORS - response = jsonify(batch) - response.headers.add("Access-Control-Allow-Origin", "*") - return response - -@app.route("/check-tokens", methods=['OPTIONS', 'POST']) -def check_tokens(): - if request.method == 'OPTIONS': - response = make_response() - # TODO: Have the interface served by the server on GET request - response.headers.add("Access-Control-Allow-Origin", "*") - response.headers.add("Access-Control-Allow-Headers", "*") - response.headers.add("Access-Control-Allow-Methods", "*") - return response - if request.method =='POST': - params = request.get_json() - text = params['text'] - tokenizer, model = generator - inputs = tokenizer([text] * 1, return_tensors="pt", truncation=True, max_length=4096).to("cuda") - # TODO: Proper CORS - response = jsonify(inputs['input_ids'][0].shape[0]) - response.headers.add("Access-Control-Allow-Origin", "*") - return response - -@app.route("/") -def index(): - return app.send_static_file("minihf.html") + prompt_node = False + evaluation_prompt = params['evaluationPrompt'] + full_prompt = context + " " + prompt + tree = TreeNode(full_prompt) + score_prompt_fn = partial(make_score_prompt_fn, evaluator) + score_prompt_fn = partial(score_prompt_fn, evaluation_prompt) + # MiniHF evaluator LoRA suffix + score_prompt_fn = partial(score_prompt_fn, "<|end|>") + # Change name to avoid overwriting global baseline evaluate_fn partial + score_fn = partial(evaluate_fn, [score_prompt_fn]) + weave_param_defaults = {"weave_n_tokens":32, "weave_budget":72, + "weave_round_budget":24, "weave_n_expand":8, + "weave_beam_width":1, "weave_max_lookahead":3, + "weave_temperature":0.25} + wp = {} + for key in weave_param_defaults.keys(): + if key in params: + try: + wp[key] = int(params[key]) + except ValueError: + wp[key] = float(params[key]) + else: + wp[key] = weave_param_defaults[key] + branches = weave_tree_search(tree=tree, + generate_fn=partial(generate_fn, + n_tokens=wp["weave_n_tokens"]), + evaluate_fn=score_fn, + budget=wp["weave_budget"], + round_budget=wp["weave_round_budget"], + n_expand=wp["weave_n_expand"], + beam_width=wp["weave_beam_width"], + max_lookahead=wp["weave_max_lookahead"], + temperature=wp["weave_temperature"]) + batch = [] + if prompt_node: + timestamp = str(time.time()) + id_ = hashlib.md5((prompt + timestamp).encode("UTF-8")).hexdigest() + batch.append({"id":id_, + "prompt":prompt, + "evaluationPrompt":evaluation_prompt, + "text":"", + "timestamp":timestamp, + "nodes":[]}) + for branch in branches: + branch_text = branch.branch_text() + timestamp = str(time.time()) + id_ = hashlib.md5((branch_text + timestamp).encode("UTF-8")).hexdigest() + batch.append({"id":id_, + "prompt": prompt, + "evaluationPrompt": evaluation_prompt, + "text":branch_text, + "timestamp":timestamp, + "nodes":branch.serialize_branch()}) + # TODO: Proper CORS + print("BATCH", batch) + response = jsonify(utils.jsonify_tensors(batch)) + response.headers.add("Access-Control-Allow-Origin", "*") + return response + + @app.route("/check-tokens", methods=['OPTIONS', 'POST']) + def check_tokens(): + if request.method == 'OPTIONS': + response = make_response() + # TODO: Have the interface served by the server on GET request + response.headers.add("Access-Control-Allow-Origin", "*") + response.headers.add("Access-Control-Allow-Headers", "*") + response.headers.add("Access-Control-Allow-Methods", "*") + return response + if request.method =='POST': + params = request.get_json() + text = params['text'] + tokenizer, model = generator + inputs = tokenizer([text] * 1, return_tensors="pt", truncation=True, max_length=4096).to(device) + # TODO: Proper CORS + response = jsonify(inputs['input_ids'][0].shape[0]) + response.headers.add("Access-Control-Allow-Origin", "*") + return response + + @app.route("/") + def index(): + return render_template('minihf.html', **config['init_weave_param']) + return app + +def main(): + parser = argparse.ArgumentParser() + parser.add_argument("--config", "-c", type=str, default="configs/jdp.json") + parser.add_argument("--device", "-d", type=str, default=utils.auto_device()) + parser.add_argument("--port", "-p", type=int, default=5000) + args = parser.parse_args() + with open(args.config, 'r') as f: + config = json.load(f) + load_models(config) + app = create_app(config, args.device) + app.run(port=args.port) + +if __name__ == "__main__": + main() diff --git a/static/minihf.html b/static/minihf.html deleted file mode 100644 index 1292d0d..0000000 --- a/static/minihf.html +++ /dev/null @@ -1,836 +0,0 @@ - - - - - - LLM Interface - - - -
-
-

Weave Generator Settings

-
- - -
- - - - - - - - - - - - - - - - -
-
- - - -
-
-
-
-
-
-
- -
-
- - - - -
- -
- 0 - - -
-
-
-
-
-
- - -
- - - diff --git a/templates/minihf.html b/templates/minihf.html new file mode 100644 index 0000000..ededa44 --- /dev/null +++ b/templates/minihf.html @@ -0,0 +1,884 @@ + + + + + + + LLM Interface + + + + +
+
+

Weave Generator Settings

+
+ + +
+ + + + + + + + + + + + + + + + +
+
+ + + +
+
+
+
+
+
+
+ +
+
+ + + + +
+ +
+ 0 + + +
+
+
+
+
+
+ + +
+ + + + \ No newline at end of file diff --git a/utils.py b/utils.py new file mode 100644 index 0000000..52bbc43 --- /dev/null +++ b/utils.py @@ -0,0 +1,24 @@ +import torch +import json + +def load_config(config_path): + with open(config_path, 'r') as f: + config = json.load(f) + return config + +def auto_device(): + if torch.cuda.is_available(): + device = torch.device('cuda') + elif torch.backends.mps.is_available(): + device = torch.device('mps') + else: + device = torch.device('cpu') + return device + +def jsonify_tensors(batch): + for result in batch: + for node in result['nodes']: + for key, value in node.items(): + if isinstance(value, torch.Tensor): + node[key] = value.item() + return batch diff --git a/weave.py b/weave.py index 1fb4f3a..10932cd 100644 --- a/weave.py +++ b/weave.py @@ -203,7 +203,7 @@ def generate_outputs(generator, text, n_tokens, n=1, batch_size=1): padding=True, truncation=True, max_length=4096 - n_tokens, - ).to("cuda") + ).to(model.device) outputs = [] with ProgressBarStreamer(total=n_tokens * n) as pbar: @@ -346,7 +346,7 @@ def evaluate_outputs(evaluator, score_prompt_fns, texts): padding=True, truncation=True, max_length=4096, - ).input_ids.to("cuda") + ).input_ids.to(model.device) logits = model(tokens).logits scores.append( torch.tensor(