transformers+huggingface训练模型

简介: 本教程介绍了如何使用 Hugging Face 的 `transformers` 库训练一个 BERT 模型进行情感分析。主要内容包括:导入必要库、下载 Yelp 评论数据集、数据预处理、模型加载与配置、定义训练参数、评估指标、实例化训练器并开始训练,最后保存模型和训练状态。整个过程详细展示了如何利用预训练模型进行微调,以适应特定任务。

[TOC]

transformers+huggingface训练模型

导入必要的库:

from datasets import load_dataset
from transformers import AutoTokenizer, AutoModelForSequenceClassification, TrainingArguments, Trainer
import numpy as np
import evaluate

CopyInsert

  • 导入 datasets 用于加载数据集。
  • 导入 transformers 中的组件,以便使用预训练的 BERT 模型和 tokenizer。
  • 导入 numpy 用于数值计算。
  • 导入 evaluate 用于计算模型预测的指标(这里是准确率)。

数据集下载:

dataset = load_dataset("yelp_review_full")

CopyInsert

  • 从 Hugging Face 的数据集中下载 Yelp 评论数据集,该数据集包含各种评论和意见。

数据预处理:

tokenizer = AutoTokenizer.from_pretrained("bert-base-cased")

def tokenize_function(examples):
    return tokenizer(examples["text"], padding="max_length", truncation=True)

CopyInsert

  • 使用预训练的 BERT model("bert-base-cased")初始化 tokenizer。
  • 定义 tokenize_function 函数,将评论文本编码成模型可接受的格式,设置填充和截断。

应用数据预处理:

tokenized_datasets = dataset.map(tokenize_function, batched=True)

CopyInsert

  • 对下载的数据集应用 tokenize_function,批量处理文本数据。

数据抽样:

small_train_dataset = tokenized_datasets["train"].shuffle(seed=42).select(range(1000))
small_eval_dataset = tokenized_datasets["test"].shuffle(seed=42).select(range(1000))

CopyInsert

  • 从训练和测试集中各随机抽取 1000 条样本,以加快训练速度和验证模型性能。

模型加载与训练配置:

model = AutoModelForSequenceClassification.from_pretrained("bert-base-cased", num_labels=5)

CopyInsert

  • 加载预训练的 BERT 模型,并指定输出标签数(5个分类)。
model_dir = "models/bert-base-cased-finetune-yelp"

training_args = TrainingArguments(
    output_dir=model_dir,
    per_device_train_batch_size=16,
    num_train_epochs=5,
    logging_steps=100
)

CopyInsert

  • 定义模型保存路径和训练参数,如每个设备的训练批大小、训练轮数和日志记录的频率。

指标评估:

metric = evaluate.load("accuracy")

def compute_metrics(eval_pred):
    logits, labels = eval_pred
    predictions = np.argmax(logits, axis=-1)
    return metric.compute(predictions=predictions, references=labels)

CopyInsert

  • 加载准确率评估指标。
  • 定义 compute_metrics 函数,通过计算预测标签和真实标签的比较来评估模型性能。

实例化 Trainer:

trainer = Trainer(
    model=model,
    args=training_args,
    train_dataset=small_train_dataset,
    eval_dataset=small_eval_dataset,
    compute_metrics=compute_metrics
)

CopyInsert

  • 创建训练器 Trainer 的实例,用于处理模型的训练过程和评估。

开始训练:

trainer.train()

CopyInsert

  • 运行训练过程。

监控 GPU 使用:

# 使用命令行工具: watch -n 1 nvidia-smi

CopyInsert

  • 提供了一个命令行工具提示,以监控 GPU 的使用情况。

保存模型和训练状态:

trainer.save_model(model_dir)
trainer.save_state()

CopyInsert

  • 保存训练完成后的模型和状态,以便后续使用。
相关文章
|
11月前
|
机器学习/深度学习 人工智能 JSON
构建AI智能体:二十八、大语言模型BERT:原理、应用结合日常场景实践全面解析
BERT是谷歌2018年推出的革命性自然语言处理模型,采用Transformer编码器架构和预训练-微调范式。其核心创新在于双向上下文理解和掩码语言建模,能有效处理一词多义和复杂语义关系。BERT通过多层自注意力机制构建深度表示,输入融合词嵌入、位置嵌入和段落嵌入,输出包含丰富上下文信息的向量。主要应用包括文本分类、命名实体识别、问答系统等,在搜索优化、智能客服、内容推荐等领域发挥重要作用。
3528 10
|
人工智能 自然语言处理 监控
大语言模型的解码策略与关键优化总结
本文系统性地阐述了大型语言模型(LLMs)中的解码策略技术原理及其应用。通过深入分析贪婪解码、束搜索、采样技术等核心方法,以及温度参数、惩罚机制等优化手段,为研究者和工程师提供了全面的技术参考。文章详细探讨了不同解码算法的工作机制、性能特征和优化方法,强调了解码策略在生成高质量、连贯且多样化文本中的关键作用。实例展示了各类解码策略的应用效果,帮助读者理解其优缺点及适用场景。
1753 20
大语言模型的解码策略与关键优化总结
|
5月前
|
人工智能 运维 前端开发
给 Hermes 装上显微镜:Agent 执行全知道
阿里云 Hermes 可观测插件基于 OpenTelemetry,追踪 Agent 推理、工具调用、Token 消耗、时延与安全风险,帮助定位成本高、响应慢、工具异常等问题。
1277 45
|
7月前
|
人工智能 自然语言处理 监控
从0到1玩转13000+OpenClaw Skill!OpenClaw阿里云/本地部署+ClawHub Skill使用攻略及避坑指南
2026年3月,OpenClaw生态迎来里程碑式更新——在云开发平台Vercel的技术支持下,官方同步上线openclaw.ai与clawhub.com两大站点。其中,clawhub.com(后统一简称为ClawHub)作为核心技能仓库,已聚合GitHub上13625个高星Skill,从PPT生成、基金操盘分析到加密货币交易、自动化运维,覆盖办公、开发、金融、创意等全场景需求,成为OpenClaw用户的“能力宝藏库”。
3166 5
|
人工智能 缓存 开发者
MCP协议究竟如何实现RAG与Agent的深度融合,打造更智能AI系统?
本文AI专家三桥君探讨了通过MCP协议实现RAG与Agent系统的深度融合,构建兼具知识理解与任务执行能力的智能系统。文章分析了传统RAG和Agent系统的局限性,提出了MCP协议的核心设计,包括标准化接口、智能缓存和动态扩展性。系统架构基于LlamaIndex和LangGraph实现服务端和客户端的协同工作,并提供了实际应用场景与生产部署指南。未来发展方向包括多模态扩展、增量更新和分布式处理等。
1214 0
|
机器学习/深度学习 人工智能 自然语言处理
编码器-解码器架构详解:Transformer如何在PyTorch中工作
本文深入解析Transformer架构,结合论文与PyTorch源码,详解编码器、解码器、位置编码及多头注意力机制的设计原理与实现细节,助你掌握大模型核心基础。建议点赞收藏,干货满满。
2383 3
|
关系型数据库 分布式数据库 数据库
阿里云数据库收费价格:MySQL、PostgreSQL、SQL Server和MariaDB引擎费用整理
阿里云数据库提供多种类型,包括关系型与NoSQL,主流如PolarDB、RDS MySQL/PostgreSQL、Redis等。价格低至21元/月起,支持按需付费与优惠套餐,适用于各类应用场景。
|
机器学习/深度学习 数据采集 自然语言处理
HuggingFace Transformers 库深度应用指南
本文首先介绍HuggingFace Tra环境配置与依赖安装,确保读者具备Python编程、机器学习和深度学习基础知识。接着深入探讨Transformers的核心组件,并通过实战案例展示其应用。随后讲解模型加载优化、批处理优化等实用技巧。在核心API部分,详细解析Tokenizers、Models、Configuration和Dataset的使用方法。文本生成章节则涵盖基础概念、GPT2生成示例及高级生成技术。最后,针对模型训练与优化,介绍预训练模型微调、超参数优化和推理加速等内容。通过这些内容,帮助读者掌握HuggingFace Transformers的深度使用,开发高效智能的NLP应用。
2411 22

热门文章

最新文章