Team Ai
Modelpublic

fifadxj/tiny-bert-sequence-classification

sourceHugging Faceapache-2.0updated 1y agoView on Hugging Face
0likes17downloads
train.py145 linesDownload Raw Back to root
1import os2 3import numpy as np4import torch5import torch.nn as nn6import transformers7from datasets import load_dataset, DatasetDict8from dotenv import load_dotenv9from transformers import AutoTokenizer, Trainer, DataCollatorWithPadding, TrainingArguments, AutoModel, \10    EarlyStoppingCallback, PreTrainedModel, AutoConfig, AutoModelForSequenceClassification, BertConfig11 12from modeling_tiny_bert_sequence_classification import TinyBertForSequenceClassification13 14 15os.environ["WANDB_DISABLED"] = "true"16 17def is_running_in_colab():18    """检查是否在Google Colab环境中运行"""19    try:20        import google.colab21        return True22    except ImportError:23        return False24 25 26is_colab = is_running_in_colab()27if is_colab:28    cache_dir = "drive/MyDrive/data/cache"29else:30    load_dotenv(".env")31    cache_dir = "cache"32 33model_checkpoint = "google-bert/bert-base-chinese"34 35config = AutoConfig.from_pretrained(model_checkpoint, cache_dir=cache_dir)36config.num_labels = 237 38model = TinyBertForSequenceClassification(config)39model.bert = AutoModel.from_pretrained(model_checkpoint, cache_dir=cache_dir)40 41print(model)42tokenizer = AutoTokenizer.from_pretrained(model_checkpoint, cache_dir=cache_dir)43 44raw_datasets = load_dataset("lansinuote/ChnSentiCorp")45# raw_datasets = DatasetDict({46#     "train": raw_datasets["train"].select(range(100)),47#     "validation": raw_datasets["validation"].select(range(100)),48#     "test": raw_datasets["test"].select(range(100)),49# })50print(raw_datasets)51 52 53def tokenize_function(example):54    return tokenizer(example["text"], truncation=True)55 56 57tokenized_datasets = raw_datasets.map(tokenize_function, batched=True)58print(tokenized_datasets)59 60tokenized_datasets.remove_columns(["text"])61tokenized_datasets.rename_column("label", "labels")62 63data_collator = DataCollatorWithPadding(tokenizer=tokenizer)64batch_size = 6465training_args = TrainingArguments(66    output_dir="tiny-bert-sequence-classification",67    learning_rate=2e-5,68    weight_decay=0.01,69    warmup_ratio=0.1,70    lr_scheduler_type="linear",71    per_device_train_batch_size=batch_size,72    per_device_eval_batch_size=batch_size,73    num_train_epochs=10,74 75    save_strategy="best",76    save_total_limit=2,77    logging_dir="./logs",78    logging_strategy="steps",79    logging_steps=len(tokenized_datasets["train"]) // batch_size,80 81    eval_strategy="epoch",82    load_best_model_at_end=True,83    metric_for_best_model="accuracy",84    greater_is_better=True,85 86    fp16=True,87    gradient_accumulation_steps=1,88    dataloader_num_workers=0,89    group_by_length=False,90    report_to=None,91    push_to_hub=False,92)93 94from sklearn.metrics import accuracy_score, precision_recall_fscore_support95def compute_metrics(eval_pred):96    """97    eval_pred 是一个 transformers.EvalPrediction 对象,包含:98        - predictions: 模型预测的 logits99        - label_ids:   真实标签100    """101    logits, labels = eval_pred102    # 如果是多分类任务103    predictions = np.argmax(logits, axis=-1)104 105    # 计算主指标106    acc = accuracy_score(labels, predictions)107    precision, recall, f1, _ = precision_recall_fscore_support(108        labels, predictions, average=None109    )110 111    # 返回 Trainer 可识别的 metrics 字典112    return {113        "accuracy": acc,114        "precision": precision[1],115        "recall": recall[1],116        "f1": f1[1],117    }118 119 120trainer = Trainer(121    model=model,122    args=training_args,123    train_dataset=tokenized_datasets["train"],124    eval_dataset=tokenized_datasets["validation"],125    data_collator=data_collator,126    processing_class=tokenizer,127    callbacks=[EarlyStoppingCallback(early_stopping_patience=2)],128    compute_metrics=compute_metrics,129)130 131trainer.train()132 133eval_results = trainer.evaluate()134print(eval_results)135 136# AutoModelForSequenceClassification.register(BertConfig, TinyBertForSequenceClassification)137# TinyBertForSequenceClassification.register_for_auto_class("AutoModel")138# TinyBertForSequenceClassification.register_for_auto_class("AutoModelForSequenceClassification")139 140model.config.auto_map = {141    "AutoModel": "modeling_tiny_bert_sequence_classification.TinyBertForSequenceClassification",142    "AutoModelForSequenceClassification": "modeling_tiny_bert_sequence_classification.TinyBertForSequenceClassification"143}144 145trainer.push_to_hub()