diff --git a/src/rblu/data_load.py b/src/rblu/data_load.py index 5b26f1f..df6051f 100644 --- a/src/rblu/data_load.py +++ b/src/rblu/data_load.py @@ -4,6 +4,7 @@ """ from typing import Optional + import matplotlib.pyplot as plt import pandas as pd from datasets import Dataset, load_dataset diff --git a/src/rblu/draw_chart/draw_metric.py b/src/rblu/draw_chart/draw_metric.py index 78421cb..dc9f587 100644 --- a/src/rblu/draw_chart/draw_metric.py +++ b/src/rblu/draw_chart/draw_metric.py @@ -1,16 +1,16 @@ import argparse -import matplotlib as mpl -import matplotlib.pyplot as plt -from pathlib import Path -from itertools import product -import os import logging +import os +from itertools import product +from pathlib import Path +import matplotlib as mpl +import matplotlib.pyplot as plt import pandas as pd import yaml from rblu.utils.name2name import translate_language, translate_model -from rblu.utils.path import CONFIG_PATH, SCORE_DIR, CHART_DIR +from rblu.utils.path import CHART_DIR, CONFIG_PATH, SCORE_DIR def _save_single_chart( diff --git a/src/rblu/draw_chart/draw_tsne_single.py b/src/rblu/draw_chart/draw_tsne_single.py index 89f071c..8c750b0 100644 --- a/src/rblu/draw_chart/draw_tsne_single.py +++ b/src/rblu/draw_chart/draw_tsne_single.py @@ -47,12 +47,12 @@ def perform_tsne(texts_flat, device): def save_tsne_results( - X_tsne, language, task, model, mode, stage, suffix="parquet" + X_tsne, language, task, model, mode, stage, suffix="parquet" ): """Save the t-SNE results to a Parquet file.""" output_df = pd.DataFrame(X_tsne, columns=["x", "y", "z"]) output_path = ( - RESULT_DIR / f"{language}_{task}_{model}_{mode}_{stage}.{suffix}" + RESULT_DIR / f"{language}_{task}_{model}_{mode}_{stage}.{suffix}" ) output_df.to_parquet(output_path, index=False) logging.info(f"Saved t-SNE results to {output_path}") @@ -73,7 +73,7 @@ def text2tsne(texts_list, language, task, model): class Tsne: def __init__( - self, language, task, model_name, mode, stage, doc_count, round + self, language, task, model_name, mode, stage, doc_count, round ) -> None: self.language = language self.task = task @@ -81,16 +81,16 @@ def __init__( self.mode = mode self.stage = stage self.path = ( - RESULT_DIR - / f"{language}_{task}_{model_name}_{mode}_{stage}.parquet" + RESULT_DIR + / f"{language}_{task}_{model_name}_{mode}_{stage}.parquet" ) self.doc_count = doc_count self.round = round def _write_and_tsne(self): data_path = ( - RESULT_DIR - / f"{self.model_name}_{self.task}_{self.stage}_{self.language}" + RESULT_DIR + / f"{self.model_name}_{self.task}_{self.stage}_{self.language}" ) texts_list = prepare_texts_for_tsne(data_path, self.mode) @@ -202,19 +202,19 @@ def _scatter_3D(round, doc_count, vector, colors, fig): def draw_tsne_single( - model_list, - language_list, - task_list, - stage, - color_family, - doc_count, - round, - suffix="html", + model_list, + language_list, + task_list, + stage, + color_family, + doc_count, + round, + suffix="html", ): """Draw t-SNE scatter plots for different models, languages, and tasks.""" for target, language in product(["q", "a"], language_list): for model_name, language, task in product( - model_list, language_list, task_list + model_list, language_list, task_list ): tsne_data = Tsne( language, task, model_name, target, stage, doc_count, round @@ -256,9 +256,9 @@ def draw_tsne_single( if args.suffix is None: args.suffix = "html" with open( - file=CONFIG_PATH, - mode="r", - encoding="utf-8", + file=CONFIG_PATH, + mode="r", + encoding="utf-8", ) as config_file: config = yaml.safe_load(config_file) # noqa: F821 draw_tsne_single( diff --git a/src/rblu/main.py b/src/rblu/main.py index 51b7cf1..a6ffd46 100755 --- a/src/rblu/main.py +++ b/src/rblu/main.py @@ -24,7 +24,7 @@ from rblu.evaluation import conservation_infer, reverse_infer, save_score from rblu.generate import APIGenerator, MyGenerator from rblu.metric import rouge_and_bert -from rblu.process.reservation_process import (get_reservation_process) +from rblu.process.reservation_process import get_reservation_process from rblu.process.reverse_process import get_reverse_process from rblu.template import apply_default_template, apply_default_zh_template from rblu.utils.path import CONFIG_PATH, RESULT_DIR, SCORE_DIR @@ -41,26 +41,31 @@ def formatTime(self, record, datefmt=None): def __init__(self): super().__init__() self.formatter_date = "%Y-%m-%d %H:%M:%S" - self._fmt = "%(asctime)s [%(thread)d] %(levelname)s - %(name)s - %(message)s" + self._fmt = ( + "%(asctime)s [%(thread)d] %(levelname)s - %(name)s - %(message)s" + ) + + formatter = IntelliJFormatter() handler = logging.FileHandler("app.log") handler.setFormatter(formatter) -logger = logging.getLogger('my.module') +logger = logging.getLogger("my.module") logger.setLevel(logging.INFO) logger.addHandler(handler) logger.propagate = False -logger.info('Initialization complete') +logger.info("Initialization complete") + def create_generator( - language: str, - model_name: str, - model_checkpoint: str | dict, - backup_mongodb: MongoCollection, - batch_size: int, - gen_kwargs: dict, - tokenizer_kwargs: dict, + language: str, + model_name: str, + model_checkpoint: str | dict, + backup_mongodb: MongoCollection, + batch_size: int, + gen_kwargs: dict, + tokenizer_kwargs: dict, ) -> APIGenerator | MyGenerator: """ Creates a generator instance based on the provided model checkpoint. @@ -86,8 +91,8 @@ def create_generator( model_checkpoint = model_checkpoint if ( - not isinstance(model_checkpoint, dict) - or model_checkpoint["type"] != "api" + not isinstance(model_checkpoint, dict) + or model_checkpoint["type"] != "api" ): return _get_local_generator( model_checkpoint, @@ -99,8 +104,8 @@ def create_generator( tokenizer_kwargs=tokenizer_kwargs, ) if ( - "key" not in model_checkpoint.keys() - or model_checkpoint["key"] == "envs" + "key" not in model_checkpoint.keys() + or model_checkpoint["key"] == "envs" ): model_checkpoint["key"] = os.getenv("GPTAPI_KEY") return APIGenerator( @@ -114,13 +119,13 @@ def create_generator( def _get_local_generator( - model_checkpoint: str, - model_name: str, - language: str, - batch_size: int, - backup_mongodb: MongoCollection, - gen_kwargs: dict, - tokenizer_kwargs: dict, + model_checkpoint: str, + model_name: str, + language: str, + batch_size: int, + backup_mongodb: MongoCollection, + gen_kwargs: dict, + tokenizer_kwargs: dict, ) -> MyGenerator: """ Initializes and returns a MyGenerator instance with the specified @@ -171,8 +176,8 @@ def _get_local_generator( def start_evaluation( - config: dict, - evaluate_task: str, + config: dict, + evaluate_task: str, ) -> None: """ Evaluates a given task using the specified configuration and process. @@ -219,8 +224,8 @@ def start_evaluation( ) output_path = ( - RESULT_DIR - / f"{model_name}_{evaluate_task}_{config['stage']}_{language}" + RESULT_DIR + / f"{model_name}_{evaluate_task}_{config['stage']}_{language}" ) if os.path.exists(output_path) and not config["force_regenerate"]: @@ -338,9 +343,9 @@ def eval(): # set the basic 'accelerate' environment on mutil-gpu write_basic_config(mixed_precision="fp16") with open( - config_path, - "r", - encoding="utf-8", + config_path, + "r", + encoding="utf-8", ) as config_file: run_config = yaml.safe_load(config_file) logging.info("Config loaded from %s", config_path) @@ -378,9 +383,9 @@ def draw(): if args.suffix is None: args.suffix = "png" with open( - file=CONFIG_PATH, - mode="r", - encoding="utf-8", + file=CONFIG_PATH, + mode="r", + encoding="utf-8", ) as config_file: config = yaml.safe_load(config_file) # noqa: F821 draw_metric( diff --git a/src/rblu/train/data_argument.py b/src/rblu/train/data_argument.py index fd2a315..f033771 100644 --- a/src/rblu/train/data_argument.py +++ b/src/rblu/train/data_argument.py @@ -1,7 +1,8 @@ """ A module to generate reverse dataset """ -from datasets import load_from_disk, Dataset + +from datasets import Dataset, load_from_disk from rblu.utils.path import RESULT_DIR @@ -13,6 +14,8 @@ dataset2train: list = [] for loop in range(5): question_column, answer_column = f"a{loop}_prompt", f"q{loop + 1}_output" - train_dataset = argued_dataset.select_columns([question_column, answer_column]) + train_dataset = argued_dataset.select_columns( + [question_column, answer_column] + ) sample_record = train_dataset.to_pandas().head(1) print(sample_record)