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("/") ) @torch.inference_mode() 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()