Team Ai
Apppublic

JetBrains-Research/commit-message-editing-visualization

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
for_labeling.py59 linesDownload Raw Back to generation_steps
1import json2 3from tqdm import tqdm4 5import config6from api_wrappers import hf_data_loader7from generation_steps import synthetic_forward8 9 10def transform(df):11    print("Generating data for labeling:")12    synthetic_forward.print_config()13    tqdm.pandas()14 15    manual_df = hf_data_loader.load_raw_rewriting_as_pandas()16 17    manual_df = manual_df.sample(frac=1, random_state=config.RANDOM_STATE).set_index(["hash", "repo"])[18        ["commit_msg_start", "commit_msg_end"]19    ]20 21    manual_df = manual_df[~manual_df.index.duplicated(keep="first")]22 23    def get_is_manually_rewritten(row):24        commit_id = (row["hash"], row["repo"])25        return commit_id in manual_df.index26 27    result = df28    result["manual_sample"] = result.progress_apply(get_is_manually_rewritten, axis=1)29 30    def get_prediction_message(row):31        commit_id = (row["hash"], row["repo"])32        if row["manual_sample"]:33            return manual_df.loc[commit_id]["commit_msg_start"]34        return row["prediction"]35 36    def get_enhanced_message(row):37        commit_id = (row["hash"], row["repo"])38        if row["manual_sample"]:39            return manual_df.loc[commit_id]["commit_msg_end"]40        return synthetic_forward.generate_end_msg(start_msg=row["prediction"], diff=row["mods"])41 42    result["enhanced"] = result.progress_apply(get_enhanced_message, axis=1)43    result["prediction"] = result.progress_apply(get_prediction_message, axis=1)44    result["mods"] = result["mods"].progress_apply(json.dumps)45 46    result.to_csv(config.DATA_FOR_LABELING_ARTIFACT)47    print("Done")48    return result49 50 51def main():52    synthetic_forward.GENERATION_ATTEMPTS = 353    df = hf_data_loader.load_full_commit_with_predictions_as_pandas()54    transform(df)55 56 57if __name__ == "__main__":58    main()59