Views
No views yet
1import joblib;
2from huggingface_hub import hf_hub_download;
3from peft import PeftModel, PeftConfig;
4from transformers import AutoTokenizer, TextClassificationPipeline, AutoModelForSequenceClassification;
5from huggingface_hub import HfApi, login
6
7# need hgf token for accessing X2BEE repo.
8with open('./api_key/HGF_TOKEN.txt', 'r') as hgf:
9 login(token=hgf.read())
10api = HfApi()
11repo_id = "x2bee/plateer_classifier_ModernBERT_v01"
12data_id = "x2bee/plateer_category_data"
13
14# Load Config, Tokenizer, Label_Encoder
15tokenizer = AutoTokenizer.from_pretrained(repo_id, subfolder="last-checkpoint")
16label_encoder_file = hf_hub_download(repo_id=data_id, repo_type="dataset", filename="label_encoder.joblib")
17label_encoder = joblib.load(label_encoder_file)
18
19# Load Model
20model = AutoModelForSequenceClassification.from_pretrained(repo_id, subfolder="last-checkpoint")
21
22import torch
23class TextClassificationPipeline(TextClassificationPipeline):
24 def __call__(self, inputs, top_k=5, **kwargs):
25 inputs = self.tokenizer(inputs, return_tensors="pt", truncation=True, padding=True, max_length=512, **kwargs)
26 inputs = {k: v.to(self.model.device) for k, v in inputs.items()}
27
28 with torch.no_grad():
29 outputs = self.model(**inputs)
30
31 probs = torch.nn.functional.softmax(outputs.logits, dim=-1)
32 scores, indices = torch.topk(probs, top_k, dim=-1)
33
34 results = []
35 for batch_idx in range(indices.shape[0]):
36 batch_results = []
37 for score, idx in zip(scores[batch_idx], indices[batch_idx]):
38 temp_list = []
39 label = self.model.config.id2label[idx.item()]
40 label = int(label.split("_")[1])
41 temp_list.append(label)
42 predicted_class = label_encoder.inverse_transform(temp_list)[0]
43
44 batch_results.append({
45 "label": label,
46 "label_decode": predicted_class,
47 "score": score.item(),
48 })
49 results.append(batch_results)
50
51 return results
52
53classifier_model = TextClassificationPipeline(tokenizer=tokenizer, model=model)
54
55def plateer_classifier(text, top_k=3):
56 result = classifier_model(text, top_k=top_k)
57 return result
58
59# run
60result = plateer_classifier("겨울 등산에서 사용할 옷")[0]
61print(result)
62
63# result
64-----------Category-----------
65{'label': 2, 'label_decode': '기능성의류', 'score': 0.9214227795600891}
66{'label': 8, 'label_decode': '스포츠', 'score': 0.07054771482944489}
67{'label': 15, 'label_decode': '패션/의류/잡화', 'score': 0.0036312134470790625}
68