Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions src/rblu/data_load.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
"""

from typing import Optional

import matplotlib.pyplot as plt
import pandas as pd
from datasets import Dataset, load_dataset
Expand Down
12 changes: 6 additions & 6 deletions src/rblu/draw_chart/draw_metric.py
Original file line number Diff line number Diff line change
@@ -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(
Expand Down
38 changes: 19 additions & 19 deletions src/rblu/draw_chart/draw_tsne_single.py
Original file line number Diff line number Diff line change
Expand Up @@ -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}")
Expand All @@ -73,24 +73,24 @@ 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
self.model_name = model_name
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)

Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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(
Expand Down
69 changes: 37 additions & 32 deletions src/rblu/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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.
Expand All @@ -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,
Expand All @@ -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(
Expand All @@ -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
Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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"]:
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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(
Expand Down
7 changes: 5 additions & 2 deletions src/rblu/train/data_argument.py
Original file line number Diff line number Diff line change
@@ -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

Expand All @@ -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)