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()