Views
No views yet
### 文章 和 ### 标签 指令格式。1import torch
2from transformers import AutoModelForCausalLM, AutoTokenizer
3
4model_name = "robertlyon/QWEN3-4B-reuters21578" # 请替换为实际模型名称
5tokenizer = AutoTokenizer.from_pretrained(model_name)
6model = AutoModelForCausalLM.from_pretrained(
7 model_name,
8 device_map="auto",
9 torch_dtype=torch.bfloat16 # 推荐bfloat16,以提高效率与性能
10)
11model.eval()
12
13def predict_topics(article: str, max_new_tokens: int = 32) -> list[str]:
14 """
15 使用微调后的Qwen3模型进行多标签分类。
16 """
17 prompt = f"### 文章\n{article.strip()}\n\n### 标签\n"
18 inputs = tokenizer(prompt, return_tensors="pt").to(model.device)
19
20 with torch.no_grad():
21 gen_ids = model.generate(
22 **inputs,
23 max_new_tokens=max_new_tokens,
24 temperature=0.1,
25 do_sample=True,
26 pad_token_id=tokenizer.eos_token_id
27 )
28
29 full_text = tokenizer.decode(gen_ids[0], skip_special_tokens=True)
30
31 if "### 标签\n" in full_text:
32 labels_text = full_text.split("### 标签\n")[-1]
33 else:
34 return []
35
36 labels = [label.strip() for label in labels_text.split(",") if label.strip()]
37 return list(dict.fromkeys(labels)) # 去重
38
39# 示例使用
40demo_text = (
41 "The U.S. Agriculture Department reported that export sales of U.S. soybeans "
42 "in the week ended Feb. 19 totaled 29.6 million bushels, compared with "
43 "47.9 million a week earlier and 19.8 million a year ago."
44)
45
46predicted_labels = predict_topics(demo_text)
47print(f"Article: '{demo_text[:70]}...'")
48print(f"Predicted Labels: {predicted_labels}")
49# 示例输出: ['soybean', 'grain', 'trade']| 指标 (Metric) | 分数 (Score) | 描述 (Description) |
|---|---|---|
| Weighted-F1 | 0.4271 | 考虑类别频率分布的F1,强调高频类别的性能表现。 |
| Micro-F1 | 0.2093 | 整体衡量的F1,受到高频标签的显著影响。 |
| Macro-F1 | 0.1835 | 不考虑标签频率的F1,强调模型在低频类别上的表现。 |
| Subset Accuracy | 0.1676 | 预测标签集合完全匹配的严格准确率。 |
| Hamming Loss | 0.0389 | 单标签误分类比例,越低越好。 |

earn、acq)上的表现相对稳健,这体现在较高的 Weighted-F1 和较低的 Hamming Loss。