hoin1218's picture
feat: add Korean benchmark models
70c7f40
Raw History Blame Contribute Delete
10.2 kB
import json
import time
from functools import lru_cache
import gradio as gr
import torch
from huggingface_hub import hf_hub_download
from sentence_transformers import SentenceTransformer, models
from transformers import AutoModel, AutoTokenizer
from benchmarking import (
BENCHMARK_TABLE_HEADERS,
DEFAULT_BENCHMARK_MODEL_IDS,
build_benchmark_row,
get_benchmark_choices,
get_benchmark_specs,
prepare_benchmark_text,
)
MODEL_REPO = "hoin1218/bge-m3-merchant-region-pair-adapter"
CONFIG_FILE = "adapter/adapter_config.json"
ADAPTER_FILE = "adapter/merchant_region_pair_projection.pt"
def get_device():
return "cuda" if torch.cuda.is_available() else "cpu"
@lru_cache(maxsize=1)
def load_model():
config_path = hf_hub_download(MODEL_REPO, CONFIG_FILE)
adapter_path = hf_hub_download(MODEL_REPO, ADAPTER_FILE)
with open(config_path, "r", encoding="utf-8") as f:
config = json.load(f)
device = get_device()
model = SentenceTransformer(config["base_model"], device=device)
dim = int(config["embedding_dimension"])
module_name = config.get("module_name", "merchant_region_pair_projection")
projection = models.Dense(
in_features=dim,
out_features=dim,
bias=True,
activation_function=None,
)
model.add_module(module_name, projection)
state = torch.load(adapter_path, map_location=device, weights_only=True)
model._modules[module_name].load_state_dict(state)
return model
@lru_cache(maxsize=8)
def load_sentence_transformer(model_id):
return SentenceTransformer(model_id, device=get_device())
class TransformersClsEncoder:
def __init__(self, model_id):
self.device = get_device()
self.tokenizer = AutoTokenizer.from_pretrained(model_id)
self.model = AutoModel.from_pretrained(model_id).to(self.device)
self.model.eval()
def encode(
self,
texts,
normalize_embeddings=True,
convert_to_numpy=True,
show_progress_bar=False,
):
del show_progress_bar
inputs = self.tokenizer(
list(texts),
padding=True,
truncation=True,
return_tensors="pt",
)
inputs = {key: value.to(self.device) for key, value in inputs.items()}
with torch.no_grad():
outputs = self.model(**inputs)
embeddings = getattr(outputs, "pooler_output", None)
if embeddings is None:
embeddings = outputs.last_hidden_state[:, 0]
if normalize_embeddings:
embeddings = torch.nn.functional.normalize(embeddings, p=2, dim=1)
if convert_to_numpy:
return embeddings.cpu().numpy()
return embeddings
@lru_cache(maxsize=3)
def load_transformers_cls_encoder(model_id):
return TransformersClsEncoder(model_id)
def load_benchmark_model(spec):
if spec.kind == "current_adapter":
return load_model()
if spec.kind == "transformers_cls":
return load_transformers_cls_encoder(spec.model_id)
return load_sentence_transformer(spec.model_id)
def compose_record(merchant, industry, region):
merchant = (merchant or "").strip()
industry = (industry or "").strip()
region = (region or "").strip()
return f"๊ฐ€๋งน์ ๋ช…: {merchant} | ์—…์ข…๋ช…: {industry} | ์ง€์—ญ: {region}"
def judge_score(score):
if score >= 0.90:
return "๊ฐ•์œ ์‚ฌ ํ›„๋ณด"
if score >= 0.50:
return "๊ฐ™์€ ์—…์ข… ํ›„๋ณด"
if score >= 0.20:
return "์œ ์‚ฌ/์• ๋งค"
return "๋‹ค๋ฅธ ์—…์ข… ๊ฐ€๋Šฅ์„ฑ ๋†’์Œ"
def encode_similarity(model, text1, text2):
embeddings = model.encode(
[text1, text2],
normalize_embeddings=True,
convert_to_numpy=True,
show_progress_bar=False,
)
return float((embeddings[0] * embeddings[1]).sum())
def compare(text1, text2):
text1 = (text1 or "").strip()
text2 = (text2 or "").strip()
if not text1 or not text2:
return 0.0, "๋‘ record๋ฅผ ๋ชจ๋‘ ์ž…๋ ฅํ•˜์„ธ์š”.", ""
model = load_model()
score = encode_similarity(model, text1, text2)
judgement = judge_score(score)
detail = (
f"score={score:.4f}\n"
f"text1={text1}\n"
f"text2={text2}"
)
return round(score, 4), judgement, detail
def compare_structured(m1, i1, r1, m2, i2, r2):
text1 = compose_record(m1, i1, r1)
text2 = compose_record(m2, i2, r2)
return (*compare(text1, text2), text1, text2)
def run_benchmark(text1, text2, selected_models):
text1 = (text1 or "").strip()
text2 = (text2 or "").strip()
if not text1 or not text2:
return [], "๋‘ record๋ฅผ ๋ชจ๋‘ ์ž…๋ ฅํ•˜์„ธ์š”."
specs = get_benchmark_specs(selected_models)
if not specs:
return [], "๋น„๊ตํ•  ๋ชจ๋ธ์„ ํ•˜๋‚˜ ์ด์ƒ ์„ ํƒํ•˜์„ธ์š”."
rows = []
details = []
for spec in specs:
started = time.perf_counter()
try:
model = load_benchmark_model(spec)
prepared1 = prepare_benchmark_text(spec, text1)
prepared2 = prepare_benchmark_text(spec, text2)
score = encode_similarity(model, prepared1, prepared2)
elapsed_ms = (time.perf_counter() - started) * 1000
judgement = judge_score(score)
rows.append(build_benchmark_row(spec, score, judgement, elapsed_ms, "OK"))
details.append(
f"{spec.label}: score={score:.4f}, elapsed_ms={elapsed_ms:.1f}"
)
except Exception as exc:
elapsed_ms = (time.perf_counter() - started) * 1000
status = f"{type(exc).__name__}: {str(exc)[:180]}"
rows.append(build_benchmark_row(spec, None, "์˜ค๋ฅ˜", elapsed_ms, status))
details.append(f"{spec.label}: ERROR {status}")
return rows, "\n".join(details)
EXAMPLES = [
[
"๊ฐ€๋งน์ ๋ช…: ์Šคํƒ€๋ฒ…์Šค ๊ฐ•๋‚จ์—ญ์  | ์—…์ข…๋ช…: ์ปคํ”ผ์ „๋ฌธ์  | ์ง€์—ญ: ์„œ์šธ ๊ฐ•๋‚จ๊ตฌ",
"๊ฐ€๋งน์ ๋ช…: ์ด๋””์•ผ์ปคํ”ผ ์—ญ์‚ผ์  | ์—…์ข…๋ช…: ์ปคํ”ผ์ „๋ฌธ์  | ์ง€์—ญ: ์„œ์šธ ๊ฐ•๋‚จ๊ตฌ",
],
[
"๊ฐ€๋งน์ ๋ช…: ์Šคํƒ€๋ฒ…์Šค ๊ฐ•๋‚จ์—ญ์  | ์—…์ข…๋ช…: ์ปคํ”ผ์ „๋ฌธ์  | ์ง€์—ญ: ์„œ์šธ ๊ฐ•๋‚จ๊ตฌ",
"๊ฐ€๋งน์ ๋ช…: ํ˜„๋Œ€์˜ค์ผ๋ฑ…ํฌ ์—ญ์‚ผ์  | ์—…์ข…๋ช…: ์ฃผ์œ ์†Œ | ์ง€์—ญ: ์„œ์šธ ๊ฐ•๋‚จ๊ตฌ",
],
[
"๊ฐ€๋งน์ ๋ช…: ์Šคํƒ€๋ฒ…์Šค ๊ฐ•๋‚จ์—ญ์  | ์—…์ข…๋ช…: ์ปคํ”ผ์ „๋ฌธ์  | ์ง€์—ญ: ์„œ์šธ ๊ฐ•๋‚จ๊ตฌ",
"๊ฐ€๋งน์ ๋ช…: ํŒŒ๋ฆฌ๋ฐ”๊ฒŒ๋œจ ์—ญ์‚ผ์  | ์—…์ข…๋ช…: ์ œ๊ณผ์  | ์ง€์—ญ: ์„œ์šธ ๊ฐ•๋‚จ๊ตฌ",
],
]
with gr.Blocks(title="Merchant Pair Similarity") as demo:
gr.Markdown("# Merchant Pair Similarity")
gr.Markdown("`๊ฐ€๋งน์ ๋ช… + ์—…์ข…๋ช… + ์ง€์—ญ` ๋‘ record๋ฅผ ๋„ฃ๊ณ  ์œ ์‚ฌ๋„ ์ ์ˆ˜๋ฅผ ํ™•์ธํ•ฉ๋‹ˆ๋‹ค.")
with gr.Tab("์ง์ ‘ ์ž…๋ ฅ"):
text1 = gr.Textbox(
label="Record 1",
lines=3,
value="๊ฐ€๋งน์ ๋ช…: ์Šคํƒ€๋ฒ…์Šค ๊ฐ•๋‚จ์—ญ์  | ์—…์ข…๋ช…: ์ปคํ”ผ์ „๋ฌธ์  | ์ง€์—ญ: ์„œ์šธ ๊ฐ•๋‚จ๊ตฌ",
)
text2 = gr.Textbox(
label="Record 2",
lines=3,
value="๊ฐ€๋งน์ ๋ช…: ์ด๋””์•ผ์ปคํ”ผ ์—ญ์‚ผ์  | ์—…์ข…๋ช…: ์ปคํ”ผ์ „๋ฌธ์  | ์ง€์—ญ: ์„œ์šธ ๊ฐ•๋‚จ๊ตฌ",
)
run = gr.Button("Compare", variant="primary")
score = gr.Number(label="Similarity", precision=4)
judgement = gr.Textbox(label="Judgement")
detail = gr.Textbox(label="Detail", lines=4)
run.click(compare, inputs=[text1, text2], outputs=[score, judgement, detail])
gr.Examples(EXAMPLES, inputs=[text1, text2], outputs=[score, judgement, detail], fn=compare)
with gr.Tab("ํ•„๋“œ ์ž…๋ ฅ"):
with gr.Row():
with gr.Column():
m1 = gr.Textbox(label="๊ฐ€๋งน์ ๋ช… 1", value="์Šคํƒ€๋ฒ…์Šค ๊ฐ•๋‚จ์—ญ์ ")
i1 = gr.Textbox(label="์—…์ข…๋ช… 1", value="์ปคํ”ผ์ „๋ฌธ์ ")
r1 = gr.Textbox(label="์ง€์—ญ 1", value="์„œ์šธ ๊ฐ•๋‚จ๊ตฌ")
with gr.Column():
m2 = gr.Textbox(label="๊ฐ€๋งน์ ๋ช… 2", value="์ด๋””์•ผ์ปคํ”ผ ์—ญ์‚ผ์ ")
i2 = gr.Textbox(label="์—…์ข…๋ช… 2", value="์ปคํ”ผ์ „๋ฌธ์ ")
r2 = gr.Textbox(label="์ง€์—ญ 2", value="์„œ์šธ ๊ฐ•๋‚จ๊ตฌ")
run_structured = gr.Button("Compare Fields", variant="primary")
score_s = gr.Number(label="Similarity", precision=4)
judgement_s = gr.Textbox(label="Judgement")
detail_s = gr.Textbox(label="Detail", lines=4)
composed1 = gr.Textbox(label="Composed Record 1")
composed2 = gr.Textbox(label="Composed Record 2")
run_structured.click(
compare_structured,
inputs=[m1, i1, r1, m2, i2, r2],
outputs=[score_s, judgement_s, detail_s, composed1, composed2],
)
with gr.Tab("๋ฒค์น˜๋งˆํฌ ๋น„๊ต"):
benchmark_text1 = gr.Textbox(
label="Record 1",
lines=3,
value="๊ฐ€๋งน์ ๋ช…: ์Šคํƒ€๋ฒ…์Šค ๊ฐ•๋‚จ์—ญ์  | ์—…์ข…๋ช…: ์ปคํ”ผ์ „๋ฌธ์  | ์ง€์—ญ: ์„œ์šธ ๊ฐ•๋‚จ๊ตฌ",
)
benchmark_text2 = gr.Textbox(
label="Record 2",
lines=3,
value="๊ฐ€๋งน์ ๋ช…: ์ด๋””์•ผ์ปคํ”ผ ์—ญ์‚ผ์  | ์—…์ข…๋ช…: ์ปคํ”ผ์ „๋ฌธ์  | ์ง€์—ญ: ์„œ์šธ ๊ฐ•๋‚จ๊ตฌ",
)
benchmark_models = gr.CheckboxGroup(
choices=get_benchmark_choices(),
value=list(DEFAULT_BENCHMARK_MODEL_IDS),
label="๋น„๊ต ๋ชจ๋ธ",
)
benchmark_run = gr.Button("Run Benchmark", variant="primary")
benchmark_results = gr.Dataframe(
headers=list(BENCHMARK_TABLE_HEADERS),
value=[],
label="Benchmark Results",
interactive=False,
)
benchmark_detail = gr.Textbox(label="Detail", lines=6)
benchmark_run.click(
run_benchmark,
inputs=[benchmark_text1, benchmark_text2, benchmark_models],
outputs=[benchmark_results, benchmark_detail],
)
gr.Markdown(
"ํŒ์ • ๊ธฐ์ค€: 0.90 ์ด์ƒ ๊ฐ•์œ ์‚ฌ ํ›„๋ณด, 0.50~0.90 ๊ฐ™์€ ์—…์ข… ํ›„๋ณด, "
"0.20~0.50 ์œ ์‚ฌ/์• ๋งค, 0.20 ๋ฏธ๋งŒ ๋‹ค๋ฅธ ์—…์ข… ๊ฐ€๋Šฅ์„ฑ ๋†’์Œ."
)
if __name__ == "__main__":
demo.launch()