userneeds1k / inference.py
bbin2022's picture
Publish optimized FP16 UserNeeds1K model
5576acd verified
Raw History Blame Contribute Delete
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("/")
)
@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()