Download app.py from hoin1218/bge-m3-merchant-region-pair-demo: direct link, hf CLI and curl.
- Browser
- Download file 10.2 kB
-
https://huggingface.co/spaces/hoin1218/bge-m3-merchant-region-pair-demo/resolve/main/app.py
- Command line
-
hf download hf://spaces/hoin1218/bge-m3-merchant-region-pair-demo/app.py
-
curl -L -o app.py https://huggingface.co/spaces/hoin1218/bge-m3-merchant-region-pair-demo/resolve/main/app.py
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" | |
| 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 | |
| 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 | |
| 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() | |