File size: 1,924 Bytes
3b2d368
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
from datasets import load_dataset
from transformers import AutoTokenizer, AutoModelForSequenceClassification, TrainingArguments, Trainer, DataCollatorWithPadding
import evaluate
import numpy as np
from transformers import EvalPrediction

task = "sst2"  # 换成 "mrpc", "rte", "qnli", "qqp", "wnli" 等

# 1. dataset & tokenizer & model
raw = load_dataset("glue", task)
tokenizer = AutoTokenizer.from_pretrained("your-model-or-tokenizer")   # 你的模型/分词器路径或名字
model = AutoModelForSequenceClassification.from_pretrained("your-model-or-checkpoint", num_labels=len(set(raw["train"]["label"])))

# 2. preprocess
def preprocess(batch):
    # 大多数 GLUE 子任务字段名是 sentence1 / sentence2
    sent1 = batch.get("sentence1") or batch.get("question") or batch.get("sentence")
    sent2 = batch.get("sentence2")
    if sent2 is None:
        return tokenizer(sent1, truncation=True)
    return tokenizer(sent1, sent2, truncation=True)

tokenized = raw.map(preprocess, batched=True)

# 3. data collator
data_collator = DataCollatorWithPadding(tokenizer)

# 4. metric
metric = evaluate.load("glue", task)

def compute_metrics(p: EvalPrediction):
    logits = p.predictions
    if isinstance(logits, tuple):  # 某些模型返回 (logits, hidden_states)
        logits = logits[0]
    preds = np.argmax(logits, axis=-1)
    return metric.compute(predictions=preds, references=p.label_ids)

# 5. trainer
training_args = TrainingArguments(
    output_dir="./out",
    per_device_train_batch_size=16,
    per_device_eval_batch_size=32,
    evaluation_strategy="epoch",
    save_strategy="epoch",
    num_train_epochs=3,
)

trainer = Trainer(
    model=model,
    args=training_args,
    train_dataset=tokenized["train"],
    eval_dataset=tokenized["validation"],
    tokenizer=tokenizer,
    data_collator=data_collator,
    compute_metrics=compute_metrics,
)

# 6. run
trainer.train()
trainer.evaluate()