backend / hf_client.py
anant-ai's picture
Update hf_client.py
52e1102 verified
Raw
History Blame Contribute Delete
4.9 kB
"""hf_client.py — Uses your newly created Model repositories"""
import os
import shutil
import time
import logging
import warnings
warnings.filterwarnings("ignore")
from huggingface_hub import login, snapshot_download
# ✅ Completely silent — no output to console
logging.disable(logging.CRITICAL)
HF_TOKEN = os.getenv("HF_TOKEN", None)
if HF_TOKEN:
login(token=HF_TOKEN, add_to_git_credential=False)
HIDE_MODEL_ERRORS = os.getenv("HIDE_MODEL_ERRORS", "true").lower() in (
"1", "true", "yes", "y",
)
PUBLIC_ERROR_MESSAGE = os.getenv(
"HF_PUBLIC_ERROR_MESSAGE",
"Model is unavailable. Please try again later.",
)
_pipelines = {
"misinfo": None,
"fakenews": None,
"emosen": None,
}
_model_ids = {
"misinfo": os.getenv("MISINFO_MODEL_ID", "anant-ai/xlm-roberta-codemix"),
"fakenews": os.getenv("FAKENEWS_MODEL_ID", "anant-ai/xlm-fakenews"),
"emosen": os.getenv("EMOSEN_MODEL_ID", "anant-ai/emo_sense"),
}
def _safe_rmtree(path: str) -> None:
try:
shutil.rmtree(path)
except Exception:
pass
def _is_corrupt_safetensors_error(err: Exception) -> bool:
msg = str(err).lower()
return (
"incomplete metadata" in msg
or "file not fully covered" in msg
or "error while deserializing header" in msg
)
def _looks_like_local_path(value: str) -> bool:
if value.startswith(("/", "./", "../", "\\\\")):
return True
return len(value) > 1 and value[1] == ":"
def _download_snapshot(model_id: str, force_download: bool = False) -> str:
if os.path.exists(model_id):
if os.path.isdir(model_id):
return model_id
raise ValueError(
"Model path points to a file. Provide the directory that contains "
"config.json, tokenizer files, and model.safetensors."
)
if _looks_like_local_path(model_id):
raise ValueError(f"Local model path not found: {model_id}")
return snapshot_download(
repo_id=model_id,
token=HF_TOKEN,
force_download=force_download,
ignore_patterns=["*.msgpack", "*.h5", "flax_model*", "tf_model*"],
)
def _build_pipeline(local_path: str):
from transformers import (
AutoTokenizer,
AutoModelForSequenceClassification,
pipeline as hf_pipeline,
)
tokenizer = AutoTokenizer.from_pretrained(
local_path,
token=HF_TOKEN,
)
model = AutoModelForSequenceClassification.from_pretrained(
local_path,
token=HF_TOKEN,
use_safetensors=True,
ignore_mismatched_sizes=True,
)
return hf_pipeline(
"text-classification",
model=model,
tokenizer=tokenizer,
top_k=None,
truncation=True,
max_length=512,
)
def _get_pipeline(name: str):
if _pipelines[name] is None:
model_id = _model_ids[name]
local_path = _download_snapshot(model_id)
try:
_pipelines[name] = _build_pipeline(local_path)
except Exception as exc:
if _is_corrupt_safetensors_error(exc):
_safe_rmtree(local_path)
local_path = _download_snapshot(model_id, force_download=True)
_pipelines[name] = _build_pipeline(local_path)
else:
raise
return _pipelines[name]
def _url_to_name(url: str) -> str:
if url in (MODEL_1_URL, MODEL_2_URL, MODEL_3_URL):
return {
MODEL_1_URL: "misinfo",
MODEL_2_URL: "fakenews",
MODEL_3_URL: "emosen",
}[url]
url = url.lower()
if "codemix" in url or "misinfo" in url:
return "misinfo"
if "fakenews" in url or "fake" in url:
return "fakenews"
if "emo" in url or "sense" in url or "sentiment" in url:
return "emosen"
return "misinfo"
def call_hf_api(url: str, text: str, token: str = "") -> dict:
name = _url_to_name(url)
for attempt in range(3):
try:
pipe = _get_pipeline(name)
result = pipe(text)
if isinstance(result, list) and len(result) > 0:
if isinstance(result[0], dict):
result = [result]
return {"status": "success", "data": result}
except Exception as e:
if attempt < 2:
time.sleep(2 ** attempt)
else:
return {
"status": "error",
"error": PUBLIC_ERROR_MESSAGE if HIDE_MODEL_ERRORS else str(e),
}
return {
"status": "error",
"error": PUBLIC_ERROR_MESSAGE if HIDE_MODEL_ERRORS else "Max retries exceeded.",
}
# Backward compatibility
MODEL_1_URL = os.getenv("MODEL_1_URL", "anant-ai/xlm-roberta-codemix")
MODEL_2_URL = os.getenv("MODEL_2_URL", "anant-ai/xlm-fakenews")
MODEL_3_URL = os.getenv("MODEL_3_URL", "anant-ai/emo_sense")