JetBrains-Research/commit-message-editing-visualization
0
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 