Team Ai
Apppublic

JetBrains-Research/commit-message-editing-visualization

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
synthetic_forward.py108 linesDownload Raw Back to generation_steps
1import pandas as pd2from tqdm import tqdm3 4import config5import dataset_statistics6from api_wrappers import grazie_wrapper7from generation_steps import examples8 9GENERATION_MULTIPLIER = 310REL_DELETIONS_THRESHOLD = 0.7511GENERATION_ATTEMPTS = 312 13 14def build_prompt(prediction, diff):15    return f"""A LLM generated a commit message for the following source code changes:16START OF THE SOURCE CODE CHANGES17{diff}18END OF THE SOURCE CODE CHANGES19 20Here is the message the LLM generated:21START OF THE COMMIT MESSAGE 22{prediction}23END OF THE COMMIT MESSAGE24 25This generated message is not perfect. Your task is to rewrite and improve it.26You have to simulate a human software developer who manually rewrites the LLM-generated commit message, 27so the message you print must share some fragments with the generated message.   28Your message should be concise. 29Follow the Conventional Commits guidelines.30Here are some examples of what you should output:31START OF THE EXAMPLES LIST32{examples.EXAMPLES_START_TO_END}33END OF THE EXAMPLES LIST34 35 36Print only the improved commit message's text after the 37token "OUTPUT".38 39OUTPUT"""40 41 42def generate_end_msg(start_msg, diff):43    prompt = build_prompt(prediction=start_msg, diff=diff)44    results = []45 46    for i in range(GENERATION_ATTEMPTS):47        end_msg_pred = grazie_wrapper.generate_for_prompt(prompt)48 49        stats = dataset_statistics.get_statistics_for_sample(50            start_msg=start_msg,51            end_msg=end_msg_pred,52        )53        if stats["deletions"] < REL_DELETIONS_THRESHOLD:54            return end_msg_pred55        else:56            results.append((stats["deletions"], end_msg_pred))57 58    results.sort()59    return results[0][1]60 61 62COLS_TO_KEEP = ["hash", "repo", "commit_msg_start", "mods", "session", "end_to_start"]63 64 65def print_config():66    print(f"NUMBER OF EXAMPLES PER PROMPT = {examples.N_EXAMPLES}")67    print(f"GENERATION_MULTIPLIER = {GENERATION_MULTIPLIER}")68    print(f"REL_DELETIONS_THRESHOLD = {REL_DELETIONS_THRESHOLD}")69    print(f"GENERATION_ATTEMPTS = {GENERATION_ATTEMPTS}")70 71 72def transform(df):73    print("Start -> send synthesis:")74    print_config()75 76    df["start_to_end"] = False77 78    generated_data = {"commit_msg_end": []}79 80    for col in COLS_TO_KEEP:81        generated_data[col] = []82 83    for _, row in tqdm(df.iterrows(), total=len(df)):84        for i in range(GENERATION_MULTIPLIER):85            commit_msg_end_pred = generate_end_msg(start_msg=row["commit_msg_start"], diff=row["mods"])86 87            generated_data["commit_msg_end"].append(commit_msg_end_pred)88            for col in COLS_TO_KEEP:89                generated_data[col].append(row[col])90 91    generated_df = pd.DataFrame.from_dict(generated_data)92    generated_df["start_to_end"] = True93 94    result = pd.concat([df, generated_df], ignore_index=True)95    result.to_csv(config.START_TO_END_ARTIFACT)96 97    print("Done")98    return result99 100 101def main():102    df = pd.read_csv(config.END_TO_START_ARTIFACT, index_col=[0])103    transform(df)104 105 106if __name__ == "__main__":107    main()108