Views
No views yet

pytorch_model.bin — веса модели (state_dict).config.json — конфигурация (input_dim, num_classes, p_dropout, classes).modeling_simple_classifier.py — определение архитектуры.vectorizer.pkl — sklearn-векторизатор (TF-IDF/Count).svd.pkl — TruncatedSVD (опционально).label_encoder.pkl — sklearn.LabelEncoder (для декодирования метки).README.md — эта карточка.1# Пример: загрузка напрямую из репозитория HF (не требует локальной копии)
2from huggingface_hub import hf_hub_download
3import json, pickle, torch
4import numpy as np
5from types import SimpleNamespace
6
7REPO = "Neweret/SimplePromptClassifier-85k"
8
9config_path = hf_hub_download(REPO, "config.json")
10weights_path = hf_hub_download(REPO, "pytorch_model.bin")
11vec_path = hf_hub_download(REPO, "vectorizer.pkl")
12svd_path = None
13try:
14 svd_path = hf_hub_download(REPO, "svd.pkl")
15except Exception:
16 svd_path = None
17le_path = hf_hub_download(REPO, "label_encoder.pkl")
18
19cfg = SimpleNamespace(**json.load(open(config_path, "r", encoding="utf-8")))
20
21# --- Динамическая модель ---
22class SimpleClassifier(torch.nn.Module):
23 def __init__(self, input_dim, num_classes, p_dropout=0.3):
24 super().__init__()
25 self.linear1 = torch.nn.Linear(input_dim, 256)
26 self.ln1 = torch.nn.LayerNorm(256)
27 self.dropout = torch.nn.Dropout(p_dropout)
28 self.linear2 = torch.nn.Linear(256, 128)
29 self.ln2 = torch.nn.LayerNorm(128)
30 self.linear_out = torch.nn.Linear(128, num_classes)
31 def forward(self, x):
32 x = torch.nn.functional.gelu(self.ln1(self.linear1(x)))
33 x = self.dropout(x)
34 x = torch.nn.functional.gelu(self.ln2(self.linear2(x)))
35 x = self.dropout(x)
36 return self.linear_out(x)
37
38model = SimpleClassifier(cfg.input_dim, cfg.num_classes, cfg.p_dropout)
39state = torch.load(weights_path, map_location="cpu")
40model.load_state_dict(state)
41model.eval()
42
43# препроцессинг
44vectorizer = pickle.load(open(vec_path, "rb"))
45svd = pickle.load(open(svd_path, "rb")) if svd_path else None
46le = pickle.load(open(le_path, "rb"))
47
48def preprocess(text):
49 X = vectorizer.transform([text])
50 if svd is not None:
51 X = svd.transform(X)
52 return X.astype(np.float32)
53
54def predict(text):
55 x = preprocess(text)
56 xb = torch.from_numpy(x).float()
57 with torch.inference_mode():
58 logits = model(xb)
59 pred = int(torch.argmax(logits, dim=1).cpu().numpy()[0])
60 return pred, le.inverse_transform([pred])[0]
61
62# пример
63print(predict("Как мне найти документацию по нашей компании?"))