Download inference.py from bbin2022/userneeds1k: direct link, hf CLI and curl.
- Browser
- Download file 6.01 kB
-
https://huggingface.co/bbin2022/userneeds1k/resolve/main/inference.py
- Command line
-
hf download hf://bbin2022/userneeds1k/inference.py
-
curl -L -o inference.py https://huggingface.co/bbin2022/userneeds1k/resolve/main/inference.py
6.01 kB
| from __future__ import annotations | |
| import argparse | |
| import json | |
| from pathlib import Path | |
| import torch | |
| from safetensors.torch import load_file | |
| from sentence_transformers import SentenceTransformer | |
| from torch import nn | |
| class MLPClassifier(nn.Module): | |
| def __init__( | |
| self, | |
| input_dim: int, | |
| hidden_dim: int, | |
| num_labels: int, | |
| dropout: float, | |
| ): | |
| super().__init__() | |
| self.network = nn.Sequential( | |
| nn.LayerNorm(input_dim), | |
| nn.Linear(input_dim, hidden_dim), | |
| nn.GELU(), | |
| nn.Dropout(dropout), | |
| nn.Linear(hidden_dim, num_labels), | |
| ) | |
| def forward(self, embeddings: torch.Tensor) -> torch.Tensor: | |
| return self.network(embeddings) | |
| class UserNeedsClassifier: | |
| def __init__( | |
| self, | |
| model_dir: str | Path = Path(__file__).parent, | |
| device: str | None = None, | |
| ): | |
| self.model_dir = Path(model_dir) | |
| self.device = torch.device( | |
| device or ("cuda" if torch.cuda.is_available() else "cpu") | |
| ) | |
| self.config = json.loads( | |
| (self.model_dir / "config.json").read_text(encoding="utf-8") | |
| ) | |
| self.category_definitions = json.loads( | |
| ( | |
| self.model_dir / self.config["category_id_definitions"] | |
| ).read_text(encoding="utf-8") | |
| ) | |
| configured_dtype = self.config.get("inference_dtype") | |
| embedding_dtype_name = self.config.get("embedding_compute_dtype") | |
| classifier_dtype_name = self.config.get("classifier_compute_dtype") | |
| dtype_by_name = { | |
| "bfloat16": torch.bfloat16, | |
| "float16": torch.float16, | |
| "float32": torch.float32, | |
| } | |
| if embedding_dtype_name is not None or classifier_dtype_name is not None: | |
| if embedding_dtype_name not in dtype_by_name: | |
| raise ValueError( | |
| f"Unsupported embedding_compute_dtype: {embedding_dtype_name}" | |
| ) | |
| if classifier_dtype_name not in dtype_by_name: | |
| raise ValueError( | |
| f"Unsupported classifier_compute_dtype: {classifier_dtype_name}" | |
| ) | |
| self.embedding_dtype = dtype_by_name[embedding_dtype_name] | |
| self.classifier_dtype = dtype_by_name[classifier_dtype_name] | |
| elif configured_dtype is None: | |
| self.embedding_dtype = ( | |
| torch.bfloat16 if self.device.type == "cuda" else torch.float32 | |
| ) | |
| self.classifier_dtype = torch.float32 | |
| else: | |
| if configured_dtype not in dtype_by_name: | |
| raise ValueError(f"Unsupported inference_dtype: {configured_dtype}") | |
| self.embedding_dtype = dtype_by_name[configured_dtype] | |
| self.classifier_dtype = dtype_by_name[configured_dtype] | |
| model_kwargs = {"dtype": self.embedding_dtype} | |
| embedding_model = Path(self.config["base_model"]) | |
| if not embedding_model.is_absolute(): | |
| embedding_model = self.model_dir / embedding_model | |
| self.embedder = SentenceTransformer( | |
| str(embedding_model), | |
| device=str(self.device), | |
| model_kwargs=model_kwargs, | |
| ) | |
| self.classifier = MLPClassifier( | |
| self.config["embedding_dim"], | |
| self.config["hidden_dim"], | |
| self.config["num_labels"], | |
| self.config["dropout"], | |
| ).to(self.device, dtype=self.classifier_dtype) | |
| state = load_file( | |
| self.model_dir / "model.safetensors", | |
| device=str(self.device), | |
| ) | |
| self.classifier.load_state_dict(state) | |
| self.classifier.eval() | |
| def category_path_to_text( | |
| self, | |
| category_id_path: str, | |
| language: str = "en", | |
| ) -> str: | |
| if language not in self.config["category_languages"]: | |
| raise ValueError(f"Unsupported category language: {language}") | |
| return " / ".join( | |
| self.category_definitions[category_id][language] | |
| for category_id in category_id_path.split("/") | |
| ) | |
| def predict( | |
| self, | |
| query: str, | |
| snippet: str, | |
| top_k: int = 5, | |
| ) -> list[dict[str, float | str]]: | |
| text = self.config["text_template"].format(query=query, snippet=snippet) | |
| encode = getattr(self.embedder, "encode_document", self.embedder.encode) | |
| embedding = encode( | |
| [text], | |
| convert_to_tensor=True, | |
| normalize_embeddings=self.config["embedding_normalized"], | |
| show_progress_bar=False, | |
| ).to(self.device, dtype=self.classifier_dtype) | |
| probabilities = self.classifier(embedding).sigmoid()[0] | |
| scores, indices = probabilities.topk(min(top_k, len(probabilities))) | |
| return [ | |
| { | |
| "category_id_path": category_id_path, | |
| "category_path_en": self.category_path_to_text(category_id_path), | |
| "score": float(score), | |
| } | |
| for score, index in zip(scores.cpu(), indices.cpu().tolist()) | |
| for category_id_path in [ | |
| self.config["id2category_path"][str(index)] | |
| ] | |
| ] | |
| def main() -> None: | |
| parser = argparse.ArgumentParser() | |
| parser.add_argument("--query", required=True) | |
| parser.add_argument("--snippet", required=True) | |
| parser.add_argument("--top-k", type=int, default=5) | |
| parser.add_argument("--model-dir", type=Path, default=Path(__file__).parent) | |
| parser.add_argument("--device") | |
| args = parser.parse_args() | |
| model = UserNeedsClassifier(args.model_dir, args.device) | |
| predictions = model.predict(args.query, args.snippet, args.top_k) | |
| print(json.dumps(predictions, ensure_ascii=False, indent=2)) | |
| if __name__ == "__main__": | |
| main() | |