Spaces:
Sleeping
Sleeping
File size: 1,609 Bytes
72e2b6e | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 | import pickle
from pathlib import Path
import scipy.sparse as sp
from sklearn.feature_extraction.text import TfidfVectorizer
def _get_tfidf_value(tfidf_config: dict, key: str):
kebab_key = key.replace("_", "-")
if key in tfidf_config:
return tfidf_config[key]
if kebab_key in tfidf_config:
return tfidf_config[kebab_key]
raise KeyError(key)
def fit_tfidf(train_texts: list[str], config: dict) -> TfidfVectorizer:
tfidf_config = config["tfidf"]
vectorizer = TfidfVectorizer(
max_features=_get_tfidf_value(tfidf_config, "max_features"),
ngram_range=tuple(_get_tfidf_value(tfidf_config, "ngram_range")),
min_df=_get_tfidf_value(tfidf_config, "min_df"),
)
vectorizer.fit(train_texts)
return vectorizer
def transform(vectorizer: TfidfVectorizer, texts: list[str]) -> sp.csr_matrix:
return vectorizer.transform(texts)
def save_vectorizer(vectorizer: TfidfVectorizer, save_path: str = "artifacts/vectorizers/tfidf.pkl") -> None:
path = Path(save_path)
path.parent.mkdir(parents=True, exist_ok=True)
with open(path, "wb") as f:
pickle.dump(vectorizer, f)
def load_vectorizer(
load_path: str = "artifacts/vectorizers/tfidf.pkl",
) -> TfidfVectorizer:
path = Path(load_path)
if not path.exists():
raise FileNotFoundError(f"Vectorizer not found: {path}")
with open(path, "rb") as f:
return pickle.load(f)
def get_feature_names(vectorizer: TfidfVectorizer) -> list[str]:
return vectorizer.get_feature_names_out().tolist()
|