Team Ai
Apppublic

ammarnasr/Code-Generation-with-Language-Specific-LoRa-Models

sourceHugging Faceopenrailupdated 3y agoView on Hugging Face
4likes
utils.py88 linesDownload Raw Back to root
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