ammarnasr/Code-Generation-with-Language-Specific-LoRa-Models
4
1 2from transformers import AutoTokenizer, AutoModelForCausalLM, GenerationConfig3import os4from peft import PeftConfig, PeftModel5import json6import jsonlines7import numpy as np8 9 10 11def initialize_tokenizer_from_huggingface(tokenizer_name):12 tokenizer = AutoTokenizer.from_pretrained(tokenizer_name)13 tokenizer.pad_token = tokenizer.eos_token14 return tokenizer15 16def initialize_causual_model_from_huffingface(model_name):17 model = AutoModelForCausalLM.from_pretrained(model_name)18 return model19 20def initialize_peft_model_from_huffingface(model_name):21 print("Loading the model from checkpoint: ", model_name, "With peft ...")22 config = PeftConfig.from_pretrained(model_name)23 model = AutoModelForCausalLM.from_pretrained(config.base_model_name_or_path)24 model = PeftModel.from_pretrained(model, model_name)25 print("Done loading the model from checkpoint: ", model_name, "With peft ...")26 model.print_trainable_parameters()27 return model28 29def initialize_generation_strategy(generation_strategy_name):30 generation_strategy = GenerationConfig.from_pretrained(generation_strategy_name)31 return generation_strategy32 33 34def stop_at_stop_token(decoded_string, stop_tokens):35 """36 Produces the prefix of decoded_string that ends at the first occurrence of37 a stop_token.38 39 WARNING: the decoded_string *must not* include the prompt, which may have stop tokens40 itself.41 """42 if stop_tokens == None:43 return decoded_string44 min_stop_index = len(decoded_string)45 for stop_token in stop_tokens:46 stop_index = decoded_string.find(stop_token)47 if stop_index != -1 and stop_index < min_stop_index:48 min_stop_index = stop_index49 return decoded_string[:min_stop_index]50 51 52 53def read_json(filename):54 with open(filename, "r") as f:55 return json.load(f)56 57 58def write_json(filename, data):59 with open(filename, "w") as f:60 json.dump(data, f, indent=4)61 62def initialize_generation_strategy_from_dict(generation_config_dict):63 generation_config = GenerationConfig(**generation_config_dict)64 return generation_config65 66 67 68def read_prompts(prompts_file_name):69 prompts = {70 "prompt_id": [],71 "prompt_text": [],72 "prompt_test": [],73 "prompt_stop_tokens": [],74 }75 with jsonlines.open(prompts_file_name) as reader:76 for prompt in reader:77 prompts["prompt_id"].append(prompt["name"])78 prompts["prompt_text"].append(prompt["prompt"])79 prompts["prompt_test"].append(prompt["tests"])80 prompts["prompt_stop_tokens"].append(prompt["stop_tokens"])81 82 promt_id_ints = [int(i.split('_')[1]) for i in prompts["prompt_id"]]83 sort_indices = np.argsort(promt_id_ints)84 85 for key in prompts:86 prompts[key] = [prompts[key][i] for i in sort_indices]87 88 return prompts