COLAS_CLASSIFIER_V2 / handler.py
YohanAcxias's picture
Upload 3 files
5164788 verified
Raw
History Blame Contribute Delete
29.4 kB
"""
HuggingFace Inference Endpoint custom handler.
Pipeline (UNIFIED semantic — ONE FAISS index covering all data sources):
1) Supervised XLM-R classifier -> top-K candidates with class probabilities
2) Bi-encoder + FAISS (single semantic) -> top-K candidates aggregated per s4_code
3) Optional CrossEncoder re-rank
4) Decision flow (same spirit as previous prod handler):
- input normalization (`uom: XX | desc\\nnon`)
- global semantic-query cleaner (newline attribute filter)
- supervised gate (with stricter gate for trash labels)
- SEMANTIC DOMINANCE override (was "catalog dominance" before)
- RRF fusion of supervised + semantic
- optional CrossEncoder bonus
- "trash label" penalty + last-mile trash-avoid
Expected repository layout (artifacts at the same level as `handler.py`):
.
├── handler.py
├── requirements.txt
├── biencoder_model_name.txt
├── rerank_model_name.txt # optional
├── faiss_semantic.index # SINGLE unified semantic index
├── semantic_table.csv # columns: s4_code, doc_text, source
└── supervised_final/
├── config.json
├── tokenizer files
├── pytorch_model.bin / model.safetensors
├── id2label.json
└── label2id.json
"""
import os
import json
import re
import unicodedata
from collections import defaultdict
from typing import Any, Dict, List, Optional, Tuple
import numpy as np
import pandas as pd
import torch
from transformers import AutoTokenizer, AutoModelForSequenceClassification
import faiss
from sentence_transformers import SentenceTransformer, CrossEncoder
# ============================================================
# Basic math utils
# ============================================================
def softmax_1d(x: np.ndarray) -> np.ndarray:
x = x - np.max(x)
e = np.exp(x)
return e / np.sum(e)
def sigmoid(x: float) -> float:
return float(1.0 / (1.0 + np.exp(-x)))
def norm_cos_to_01(cos_sim: float) -> float:
v = (cos_sim + 1.0) / 2.0
return float(max(0.0, min(1.0, v)))
def faiss_score_to_01(score_raw: float, metric_type: Optional[int]) -> float:
if metric_type == faiss.METRIC_L2:
return float(np.exp(-float(score_raw)))
s = float(score_raw)
if -1.01 <= s <= 1.01:
return norm_cos_to_01(s)
return float(sigmoid(s / 5.0))
def parse_bool(v: Any, default: bool = True) -> bool:
if v is None:
return default
if isinstance(v, bool):
return v
if isinstance(v, (int, float)):
return bool(v)
if isinstance(v, str):
return v.strip().lower() in ("1", "true", "yes", "y", "on")
return bool(v)
# ============================================================
# Cleaning helpers
# ============================================================
def strip_after_first_bar_keep_newlines(s: str) -> str:
if not s or not isinstance(s, str):
return ""
if "|" in s:
return s.split("|", 1)[1]
return s
def fix_encoding(text: str) -> str:
if not text or not isinstance(text, str):
return ""
t = text
try:
t2 = t.encode("latin1", errors="strict").decode("utf-8", errors="strict")
if "Ã" in t or "�" in t:
return t2
except Exception:
pass
try:
t2 = t.encode("cp1252", errors="strict").decode("utf-8", errors="strict")
if "Ã" in t or "�" in t:
return t2
except Exception:
pass
return t
def clean_custom_code_name(name: str) -> str:
if not name or not isinstance(name, str):
return ""
name = fix_encoding(name)
name = re.sub(r"\(\d{5,}\)", "", name)
name = re.sub(r"^NPO\s*-\s*\d+\s*-\s*", "", name, flags=re.IGNORECASE)
name = re.sub(r"^\d+\s*[-:]\s*", "", name)
name = re.sub(r"[()]", " ", name)
name = re.sub(r"[\[\]]", " ", name)
name = re.sub(r"[,;/\\|#@&*=+]", " ", name)
name = re.sub(r"[-–—]", " ", name)
name = re.sub(r"\s+", " ", name).strip()
name = unicodedata.normalize("NFC", name)
return name
# ============================================================
# GLOBAL semantic query cleaner with newline usefulness check
# ============================================================
SEM_RE_MULTI_SPACE = re.compile(r"\s+")
SEM_RE_BRACKET_BLOCK = re.compile(r"\[[^\]]{1,60}\]")
SEM_RE_TRAILING_PARENS_CODELIKE = re.compile(r"\s*\(\s*[A-Za-z0-9][A-Za-z0-9._/\-]{3,60}\s*\)\s*$")
SEM_RE_JUNK_CODES = re.compile(
r"(?i)\b(ref|réf|reference|référence|code|article|art|sku|id)\s*[:#]?\s*[A-Za-z0-9][A-Za-z0-9./_-]{2,}\b"
)
SEM_RE_UNITS = re.compile(
r"(?i)\b("
r"kg|kilogramme?s?|g|gr|t|tonne?s?|"
r"mm|millim[eè]tre?s?|cm|centim[eè]tre?s?|m\b|m[èe]tre?s?|"
r"l\b|litre?s?|ml|cl|"
r"bar|pfa|pn|psi|pa|mpa|"
r"w\b|kw|v\b|a\b|ah|hz|"
r"°c|c°"
r")\b"
)
SEM_ATTR_KEYS = [
"poids", "poid", "dimension", "dimensions", "dim", "longueur", "largeur", "hauteur",
"epaisseur", "épaisseur", "diametre", "diamètre", "ø", "diam",
"référence", "reference", "ref", "réf", "code", "article", "sku", "id",
"forme", "usage", "matiere", "matière", "type", "couleur", "coloris",
"pfa", "pn", "pression", "pressure", "température", "temperature",
"conditionnement", "colisage", "unité", "uom", "unite",
]
SEM_RE_ATTR_KV = re.compile(
r"(?i)\b(" + "|".join(re.escape(k) for k in SEM_ATTR_KEYS) + r")\b\s*[:=]"
)
SEM_RE_MANY_KV = re.compile(r".*[:=].*[:=].*")
SEM_RE_JUST_REFERENCE = re.compile(r"(?i)^\s*(référence|reference|ref|réf)\b")
def _digit_ratio(s: str) -> float:
digits = sum(ch.isdigit() for ch in s)
letters = sum(ch.isalpha() for ch in s)
denom = max(1, digits + letters)
return digits / denom
def _looks_like_attribute_line(line: str) -> bool:
if not line:
return True
l = line.strip()
low = l.lower()
if len(l) <= 6 and SEM_RE_UNITS.search(l):
return True
if SEM_RE_JUST_REFERENCE.match(l):
return True
if SEM_RE_ATTR_KV.search(l):
return True
if SEM_RE_MANY_KV.match(l):
return True
if _digit_ratio(l) >= 0.22 and SEM_RE_UNITS.search(l):
return True
if l.count(" - ") >= 2 and (":" in l or "=" in l):
return True
if _digit_ratio(l) >= 0.35:
return True
if " pfa" in (" " + low + " ") or " pn" in (" " + low + " "):
return True
return False
def _select_useful_newline_parts(raw: str, max_extra_lines: int = 1) -> str:
t = (raw or "").replace("\r", "\n")
t = re.sub(r"\n\s*(oui|non)\s*$", "", t, flags=re.IGNORECASE).strip()
lines = [ln.strip() for ln in t.split("\n") if ln.strip()]
if not lines:
return ""
kept = [lines[0]]
extras = 0
for ln in lines[1:]:
if _looks_like_attribute_line(ln):
continue
if len(re.findall(r"[A-Za-zÀ-ÿ]{3,}", ln)) < 2:
continue
kept.append(ln)
extras += 1
if extras >= max_extra_lines:
break
return " ".join(kept)
def clean_semantic_query(
raw_text: str,
*,
drop_leading_long_number: bool = False,
max_extra_lines: int = 1,
) -> str:
if not raw_text or not isinstance(raw_text, str):
return ""
t = fix_encoding(raw_text)
t = strip_after_first_bar_keep_newlines(t)
t = _select_useful_newline_parts(t, max_extra_lines=max_extra_lines)
t = SEM_RE_BRACKET_BLOCK.sub(" ", t)
if drop_leading_long_number:
t = re.sub(r"^\d{4,}\s+", "", t.strip())
prev = None
while prev != t:
prev = t
t = SEM_RE_TRAILING_PARENS_CODELIKE.sub("", t).strip()
t = SEM_RE_JUNK_CODES.sub(" ", t)
t = SEM_RE_MULTI_SPACE.sub(" ", t).strip()
return t
# ============================================================
# Handler
# ============================================================
class EndpointHandler:
"""
HF Inference Endpoint Custom Handler — UNIFIED semantic (single FAISS index).
Artifacts expected in repo root:
- supervised_final/ (HF transformers classifier)
- faiss_semantic.index + semantic_table.csv (SINGLE semantic index)
- biencoder_model_name.txt
- optional: rerank_model_name.txt
"""
def __init__(self, path: str = ""):
self.device = "cuda" if torch.cuda.is_available() else "cpu"
# ---- 1) Supervised classifier ----
self.sup_dir = os.path.join(path, "supervised_final")
self.tokenizer = AutoTokenizer.from_pretrained(self.sup_dir, use_fast=True)
self.model = AutoModelForSequenceClassification.from_pretrained(self.sup_dir).to(self.device)
self.model.eval()
id2label_path = os.path.join(self.sup_dir, "id2label.json")
if os.path.exists(id2label_path):
with open(id2label_path, "r", encoding="utf-8") as f:
self.id2label = {int(k): v for k, v in json.load(f).items()}
else:
self.id2label = {int(k): v for k, v in self.model.config.id2label.items()}
# ---- 2) Bi-encoder ----
self.biencoder = None
biencoder_name_path = os.path.join(path, "biencoder_model_name.txt")
if os.path.exists(biencoder_name_path):
with open(biencoder_name_path, "r", encoding="utf-8") as f:
biencoder_name = f.read().strip()
if biencoder_name:
self.biencoder = SentenceTransformer(biencoder_name, device=self.device)
# ---- 3) UNIFIED semantic FAISS ----
self.sem_index = None
self.sem_metric = None
self.sem_table = None
self.sem_text_col = None
sem_index_path = os.path.join(path, "faiss_semantic.index")
sem_table_path = os.path.join(path, "semantic_table.csv")
if os.path.exists(sem_index_path) and os.path.exists(sem_table_path):
self.sem_index = faiss.read_index(sem_index_path)
self.sem_metric = getattr(self.sem_index, "metric_type", None)
self.sem_table = pd.read_csv(sem_table_path)
for c in ["doc_text", "label_text", "text", "description", "name", "title"]:
if c in self.sem_table.columns:
self.sem_text_col = c
break
# ---- 4) Optional reranker ----
self.reranker = None
rerank_name_path = os.path.join(path, "rerank_model_name.txt")
if os.path.exists(rerank_name_path):
with open(rerank_name_path, "r", encoding="utf-8") as f:
rerank_name = f.read().strip()
if rerank_name:
self.reranker = CrossEncoder(rerank_name, max_length=512, device=self.device)
# -----------------------------
# Input normalization / query formatting
# -----------------------------
def _normalize_incoming(self, raw_text: str, uom: Optional[str]) -> Tuple[str, str]:
t = (raw_text or "").strip()
m = re.match(r"^\s*uom\s*:\s*([A-Za-z0-9]+)\s*\|\s*(.*)$", t, flags=re.IGNORECASE | re.DOTALL)
if m:
uom = uom or m.group(1).strip().upper()
t = m.group(2).strip()
t = re.sub(r"\n\s*(oui|non)\s*$", "", t, flags=re.IGNORECASE).strip()
t = t.replace("\r", "\n")
t = re.sub(r"[ \t]+", " ", t).strip()
final_uom = (uom or "UN").strip().upper()
return t, final_uom
def _make_query_text(self, desc: str, uom: str) -> str:
return f"[UOM={uom}] {(desc or '').strip()}".strip()
def _make_semantic_query(self, raw_text: str, *, drop_leading_long_number: bool, max_extra_lines: int) -> str:
return clean_semantic_query(
raw_text,
drop_leading_long_number=drop_leading_long_number,
max_extra_lines=max_extra_lines,
)
# -----------------------------
# Supervised topK
# -----------------------------
def _supervised_topk(self, text: str, top_k: int = 20) -> List[Dict[str, Any]]:
inputs = self.tokenizer(
text, truncation=True, padding=True, max_length=256, return_tensors="pt",
).to(self.device)
with torch.no_grad():
logits = self.model(**inputs).logits[0].detach().cpu().numpy()
probs = softmax_1d(logits)
idx = np.argsort(-probs)[:top_k]
return [
{"label_id": int(i), "s4_code": str(self.id2label.get(int(i), str(i))), "score": float(probs[i])}
for i in idx
]
# -----------------------------
# Semantic retrieval (single unified index, max-pool per s4_code)
# -----------------------------
def _semantic_topk(
self,
raw_query: str,
*,
top_k: int,
semantic_drop_leading_long_number: bool,
semantic_max_extra_lines: int,
faiss_overscan: int = 4,
) -> List[Dict[str, Any]]:
if self.sem_index is None or self.sem_table is None or self.biencoder is None:
return []
if "s4_code" not in self.sem_table.columns:
return []
query = self._make_semantic_query(
raw_query,
drop_leading_long_number=semantic_drop_leading_long_number,
max_extra_lines=semantic_max_extra_lines,
)
if not query:
return []
emb = self.biencoder.encode([query], normalize_embeddings=True)
emb = np.asarray(emb, dtype=np.float32)
k_raw = min(int(max(top_k, 1) * max(faiss_overscan, 1)), self.sem_index.ntotal)
scores, idx = self.sem_index.search(emb, max(k_raw, top_k))
scores = scores[0].tolist()
idx = idx[0].tolist()
best_per_code: Dict[str, Dict[str, Any]] = {}
for s, i in zip(scores, idx):
if i < 0 or i >= len(self.sem_table):
continue
row = self.sem_table.iloc[i]
code = str(row.get("s4_code", "")).strip()
if not code:
continue
txt = ""
if self.sem_text_col and self.sem_text_col in self.sem_table.columns:
txt = str(row.get(self.sem_text_col, "") or "")
txt_one_line = txt.replace("\r", " ").replace("\n", " ").strip()
src = str(row["source"]) if "source" in self.sem_table.columns else "semantic"
txt_clean = clean_semantic_query(
txt_one_line, drop_leading_long_number=False, max_extra_lines=0
) or txt_one_line
cand = {
"source": src,
"faiss_id": int(i),
"s4_code": code,
"label_text": txt_one_line,
"label_text_clean": txt_clean,
"score_raw": float(s),
"score_01": float(faiss_score_to_01(float(s), self.sem_metric)),
"semantic_query": query,
}
prev = best_per_code.get(code)
if prev is None or cand["score_raw"] > prev["score_raw"]:
best_per_code[code] = cand
return sorted(best_per_code.values(), key=lambda x: x.get("score_01", 0.0), reverse=True)[:top_k]
# -----------------------------
# Fusion + optional rerank (RRF rank-based)
# -----------------------------
def _fuse_and_choose(
self,
raw_query: str,
sup: List[Dict[str, Any]],
sem: List[Dict[str, Any]],
*,
use_rerank: bool,
trash_labels: set,
trash_penalty_rrf: float,
weights: Tuple[float, float, float] = (1.0, 1.0, 0.5), # (sup, sem, rerank)
rerank_topn: int = 40,
debug_topn: int = 10,
rrf_k: int = 60,
semantic_drop_leading_long_number: bool = False,
semantic_max_extra_lines: int = 1,
) -> Dict[str, Any]:
w_sup, w_sem, w_rerank = weights
def rrf(rank: Optional[int]) -> float:
if rank is None:
return 0.0
return 1.0 / (rrf_k + rank)
def build_rrf_sum_map(items: List[Dict[str, Any]], code_key="s4_code"):
score_map = defaultdict(float)
first_rank = {}
for rank, x in enumerate(items, start=1):
code = str(x.get(code_key) or "")
if not code:
continue
score_map[code] += rrf(rank)
first_rank.setdefault(code, rank)
return score_map, first_rank
sup_score_map, sup_first_rank = build_rrf_sum_map(sup)
sem_score_map, sem_first_rank = build_rrf_sum_map(sem)
text_map: Dict[str, str] = {}
for x in sem:
code = str(x["s4_code"])
text_map[code] = x.get("label_text_clean") or x.get("label_text") or code
candidates = set(sup_score_map.keys()) | set(sem_score_map.keys())
if not candidates:
return {"prediction": None, "method": "no_candidates", "fused_topk": []}
fused = []
for code in candidates:
score = (
w_sup * sup_score_map.get(code, 0.0)
+ w_sem * sem_score_map.get(code, 0.0)
)
if code in trash_labels:
score -= float(trash_penalty_rrf)
fused.append({
"s4_code": code,
"score": float(score),
"components": {
"sup_rank": sup_first_rank.get(code),
"sem_rank": sem_first_rank.get(code),
},
"text_for_rerank": text_map.get(code, code),
})
fused.sort(key=lambda x: x["score"], reverse=True)
rerank_query = self._make_semantic_query(
raw_query,
drop_leading_long_number=semantic_drop_leading_long_number,
max_extra_lines=semantic_max_extra_lines,
)
if not rerank_query:
rerank_query = re.sub(r"\s+", " ", (raw_query or "").replace("\r", " ").replace("\n", " ")).strip()
if use_rerank and self.reranker is not None and fused:
topn = fused[: min(rerank_topn, len(fused))]
pairs = [[rerank_query, c["text_for_rerank"]] for c in topn]
rr_raw = self.reranker.predict(pairs).tolist()
order = np.argsort(-np.array(rr_raw))
rr_rank_map = {int(idx): r for r, idx in enumerate(order, start=1)}
for i, c in enumerate(topn):
c["rerank_raw"] = float(rr_raw[i])
c["rerank_rank"] = rr_rank_map[i]
c["score"] = float(c["score"] + w_rerank * (1.0 / (rrf_k + c["rerank_rank"])))
fused[: len(topn)] = topn
fused.sort(key=lambda x: x["score"], reverse=True)
return {"prediction": fused[0], "method": "fused_rrf_rank", "fused_topk": fused[:debug_topn]}
# -----------------------------
# HF entrypoint
# -----------------------------
def __call__(self, data: Dict[str, Any]) -> Dict[str, Any]:
inputs = data.get("inputs")
if inputs is None:
return {"error": "Missing 'inputs'."}
top_k = int(data.get("top_k", 20))
sem_top_k = int(data.get("sem_top_k", 50))
# Supervised gating
sup_min_for_fusion = float(data.get("sup_min_for_fusion", data.get("threshold", 0.80)))
sup_margin_min = float(data.get("sup_margin_min", 0.03))
sup_fusion_topn = int(data.get("sup_fusion_topn", 5))
# Semantic cleaning knobs
semantic_max_extra_lines = int(data.get("semantic_max_extra_lines", 1))
semantic_drop_leading_long_number = parse_bool(
data.get("semantic_drop_leading_long_number", False), default=False
)
# Semantic dominance knobs (replaces former 'catalog dominance')
sem_dom_min = float(data.get("sem_dom_min", data.get("cat_dom_min", 0.83)))
sem_dom_gap = float(data.get("sem_dom_gap", data.get("cat_dom_gap", 0.02)))
sem_dom_rerank_topn = int(data.get("sem_dom_rerank_topn", data.get("cat_dom_rerank_topn", 10)))
use_rerank = parse_bool(data.get("use_rerank", True), default=True)
# Trash labels handling
trash_labels = set(data.get("trash_labels", ["AP0009"]))
trash_penalty_rrf = float(data.get("trash_penalty_rrf", 0.02))
trash_sup_min_for_fusion = float(data.get("trash_sup_min_for_fusion", 0.98))
trash_sup_margin_min = float(data.get("trash_sup_margin_min", 0.15))
global_uom = data.get("uom") or data.get("CODE_UNITE") or data.get("code_unite")
common_kwargs = dict(
top_k=top_k,
sem_top_k=sem_top_k,
use_rerank=use_rerank,
sup_min_for_fusion=sup_min_for_fusion,
sup_margin_min=sup_margin_min,
sup_fusion_topn=sup_fusion_topn,
sem_dom_min=sem_dom_min,
sem_dom_gap=sem_dom_gap,
sem_dom_rerank_topn=sem_dom_rerank_topn,
trash_labels=trash_labels,
trash_penalty_rrf=trash_penalty_rrf,
trash_sup_min_for_fusion=trash_sup_min_for_fusion,
trash_sup_margin_min=trash_sup_margin_min,
semantic_max_extra_lines=semantic_max_extra_lines,
semantic_drop_leading_long_number=semantic_drop_leading_long_number,
)
def run_one(x: Any) -> Dict[str, Any]:
if isinstance(x, str):
return self._predict_one(text=x, uom=global_uom, **common_kwargs)
if isinstance(x, dict):
text = x.get("text") or x.get("description") or x.get("inputs") or ""
uom = x.get("uom") or x.get("CODE_UNITE") or x.get("code_unite") or global_uom
return self._predict_one(text=text, uom=uom, **common_kwargs)
return {"error": "Each item in inputs must be str or dict."}
if isinstance(inputs, (str, dict)):
return run_one(inputs)
if isinstance(inputs, list):
return {"results": [run_one(x) for x in inputs]}
return {"error": "'inputs' must be a string, dict, or a list of strings/dicts."}
def _predict_one(
self,
*,
text: str,
uom: Optional[str],
top_k: int,
sem_top_k: int,
use_rerank: bool,
sup_min_for_fusion: float,
sup_margin_min: float,
sup_fusion_topn: int,
sem_dom_min: float,
sem_dom_gap: float,
sem_dom_rerank_topn: int,
trash_labels: set,
trash_penalty_rrf: float,
trash_sup_min_for_fusion: float,
trash_sup_margin_min: float,
semantic_max_extra_lines: int,
semantic_drop_leading_long_number: bool,
) -> Dict[str, Any]:
desc, final_uom = self._normalize_incoming(text, uom=uom)
clf_query = self._make_query_text(desc, uom=final_uom)
sup_all = self._supervised_topk(clf_query, top_k=top_k)
bert_top1 = float(sup_all[0]["score"]) if sup_all else 0.0
bert_top2 = float(sup_all[1]["score"]) if len(sup_all) > 1 else 0.0
bert_margin = bert_top1 - bert_top2
bert_top1_code = str(sup_all[0]["s4_code"]) if sup_all else ""
if bert_top1_code in trash_labels:
bert_used_in_fusion = (bert_top1 >= trash_sup_min_for_fusion) and (bert_margin >= trash_sup_margin_min)
else:
bert_used_in_fusion = (bert_top1 >= sup_min_for_fusion) and (bert_margin >= sup_margin_min)
sup_for_fusion = sup_all[: max(0, int(sup_fusion_topn))] if bert_used_in_fusion else []
raw_sem_query = text if isinstance(text, str) else str(text)
sem = self._semantic_topk(
raw_query=raw_sem_query,
top_k=sem_top_k,
semantic_drop_leading_long_number=semantic_drop_leading_long_number,
semantic_max_extra_lines=semantic_max_extra_lines,
)
sem_query_used = self._make_semantic_query(
raw_sem_query,
drop_leading_long_number=semantic_drop_leading_long_number,
max_extra_lines=semantic_max_extra_lines,
)
# -----------------------------
# SEMANTIC DOMINANCE override (was catalog_dominance)
# -----------------------------
if sem:
top1 = sem[0]
top1_code = str(top1["s4_code"])
top1_s = float(top1.get("score_01", 0.0))
top2_s = 0.0
for x in sem[1:]:
if str(x["s4_code"]) != top1_code:
top2_s = float(x.get("score_01", 0.0))
break
if (top1_code not in trash_labels) and (top1_s >= sem_dom_min) and ((top1_s - top2_s) >= sem_dom_gap):
if use_rerank and self.reranker is not None:
rerank_query = sem_query_used or re.sub(r"\s+", " ", raw_sem_query).strip()
unique = []
seen = set()
for x in sem[: max(1, int(sem_dom_rerank_topn))]:
code = str(x["s4_code"])
if code in seen:
continue
seen.add(code)
unique.append(x)
pairs = [
[rerank_query, (u.get("label_text_clean") or u.get("label_text") or str(u["s4_code"]))]
for u in unique
]
rr_raw = self.reranker.predict(pairs).tolist()
best_idx = int(np.argmax(np.array(rr_raw)))
best = unique[best_idx]
return {
"input": text,
"uom": final_uom,
"prediction": {
"s4_code": str(best["s4_code"]),
"score": float(best.get("score_01", 0.0)),
"components": {"semantic_dom": float(best.get("score_01", 0.0))},
"text_for_rerank": best.get("label_text_clean") or best.get("label_text"),
"rerank_raw": float(rr_raw[best_idx]),
"source": best.get("source"),
},
"method": "semantic_dominance_rerank",
"supervised_topk": sup_all[:10],
"semantic_topk": sem[:10],
"fused_topk": [],
"bert_used_in_fusion": False,
"bert_top1": bert_top1,
"bert_margin": bert_margin,
"semantic_query_used": sem_query_used,
}
return {
"input": text,
"uom": final_uom,
"prediction": {
"s4_code": top1_code,
"score": top1_s,
"components": {"semantic_dom": top1_s},
"text_for_rerank": top1.get("label_text_clean") or top1.get("label_text"),
"source": top1.get("source"),
},
"method": "semantic_dominance",
"supervised_topk": sup_all[:10],
"semantic_topk": sem[:10],
"fused_topk": [],
"bert_used_in_fusion": False,
"bert_top1": bert_top1,
"bert_margin": bert_margin,
"semantic_query_used": sem_query_used,
}
# -----------------------------
# Fusion (RRF) with trash penalty
# -----------------------------
fused = self._fuse_and_choose(
raw_query=raw_sem_query,
sup=sup_for_fusion,
sem=sem,
use_rerank=use_rerank,
trash_labels=trash_labels,
trash_penalty_rrf=trash_penalty_rrf,
weights=(1.0, 1.0, 0.5),
rerank_topn=40,
debug_topn=10,
semantic_drop_leading_long_number=semantic_drop_leading_long_number,
semantic_max_extra_lines=semantic_max_extra_lines,
)
pred = fused.get("prediction")
# last-mile trash avoid if close tie
if pred and str(pred.get("s4_code", "")) in trash_labels:
for alt in fused.get("fused_topk", [])[1:]:
if str(alt.get("s4_code", "")) not in trash_labels:
if float(pred.get("score", 0.0)) - float(alt.get("score", 0.0)) <= 0.01:
pred = alt
fused["method"] = "fused_rrf_rank_trash_avoid"
break
return {
"input": text,
"uom": final_uom,
"prediction": pred,
"method": fused.get("method"),
"supervised_topk": sup_all[:10],
"semantic_topk": sem[:10],
"fused_topk": fused.get("fused_topk", []),
"bert_used_in_fusion": bert_used_in_fusion,
"bert_top1": bert_top1,
"bert_margin": bert_margin,
"semantic_query_used": sem_query_used,
"semantic_max_extra_lines": int(semantic_max_extra_lines),
"semantic_drop_leading_long_number": bool(semantic_drop_leading_long_number),
"trash_labels": sorted(list(trash_labels)),
}