基础信息说明

  • 本文以Seq2SeqTrainer作为实例,来讨论其模型训练时的数据加载方式
  • 预训练模型:opus-mt-en-zh
  • 数据集:本地数据集
  • 任务:en-zh 机器翻译

数据加载

Trainer的数据加载方式主要分为两种:基于torch.utils.data.Dataset的方式加载 和 基于huggingface自带的Datasets的方式(下文用huggingface / Datasets表示)加载。以下是一些需要注意的点:(1)Seq2SeqTrainer()的train_dataset和eval_dataset参数的所传实参应为字典类型;(2)该字典实参的keys应当覆盖模型运行所需要的数据参数(本文需要包括的有:'input_ids', 'attention_mask', 'labels');(3)使用huggingface / Datasets方法加载时,传给train_dataset和eval_dataset的字典实参中,多余的key(未在模型运行所需输入参数列表中)及其相关数据数,将会在训练之前被剔除。

torch.utils.data.Dataset

重载Dataset类(dataset.py)

# -*- coding: utf-8 -*-
from torch.utils.data import Dataset class CDNDataset(Dataset):
def __init__(self, samples):
super(CDNDataset, self).__init__()
self.samples = samples def __getitem__(self, ite):
res = {k_: v_[ite]for k_, v_ in self.samples.items()}
return res def __len__(self):
return len(self.samples['labels'])

加载引用(main.py 后文代码同属于本文件)

from transformers import AutoTokenizer, DataCollatorForSeq2Seq, AutoModelForSeq2SeqLM, Seq2SeqTrainingArguments, Seq2SeqTrainer
from dataset import CDNDataset

读取数据

# 读取训练集
with open('raw_data/txt_en.txt', 'r', encoding='utf-8') as fr_en, open('raw_data/txt_zh.txt', 'r', encoding='utf-8') as fr_zh:
train_data = tokenizer([str_.strip() for str_ in fr_en.readlines()], max_length=128, padding=True,truncation=True)
# 将tokenized的中文序列对应的input_ids作为输入数据的标签
train_data['labels'] = tokenizer([str_.strip() for str_ in fr_zh.readlines()], max_length=128,
padding=True,truncation=True)["input_ids"]
fr_en.close()
fr_zh.close()
train_data = CDNDataset(train_data) # 读取验证集
with open('raw_data/test_txt_en.txt', 'r', encoding='utf-8') as fr_en, open('raw_data/test_txt_zh.txt', 'r', encoding='utf-8') as fr_zh:
dev_data = tokenizer([str_.strip() for str_ in fr_en.readlines()], max_length=128, padding=True,truncation=True)
dev_data['labels'] = tokenizer([str_.strip() for str_ in fr_zh.readlines()], max_length=128,
padding=True,truncation=True)["input_ids"]
fr_en.close()
fr_zh.close()
dev_data = CDNDataset(dev_data)

huggingface / Datasets

修改main.py中数据集读取部分的代码

from datasets import load_dataset
from transformers import AutoTokenizer, DataCollatorForSeq2Seq, AutoModelForSeq2SeqLM, Seq2SeqTrainingArguments, Seq2SeqTrainer
"""利用load_dataset()来读取数据:
- 该方法支持.txt、.csv、.json等文件格式
- 返回结果是一个字典类型
- 读取.txt文件时,若不指定名称,这key为"text", 且会返回文本中的样本数(段落数)
- 在读取.json文件时,若所有样本放在一个josn文件中,则返回的样本数为1(无法优雅地调用train_test_split()进行数据集分割),名称为默认名或者最层字典所 对应的keys;
- 将每个json文件仅存放一个样本,并把这些文件放在某一目录,可使利用load_dataset()正确计算出样本数。但该目录下每个.json文件命名风格要一致(例如:txt1.json、txt2.json、、、),文件名差异较大的话,系统会只读取某一类命名格式相近的文件中的数据。
"""
books = load_dataset("raw_data", data_dir='test_en', name='translation') books = books["train"].train_test_split(test_size=0.15) source_lang = "en"
target_lang = "zh"
prefix = "translate English to Chinese: " # 其实我也还没搞懂为啥要加这样一个前缀 def preprocess_function(examples):
inputs = [prefix + example[source_lang] for example in examples["translation"]]
targets = [example[target_lang] for example in examples["translation"]]
model_inputs = tokenizer(inputs, max_length=128, truncation=True) with tokenizer.as_target_tokenizer():
labels = tokenizer(targets, max_length=128, truncation=True) model_inputs["labels"] = labels["input_ids"]
return model_inputs tokenized_books = books.map(preprocess_function, batched=True)

模型及参数加载

tokenizer = AutoTokenizer.from_pretrained("opus-mt-en-zh")
model = AutoModelForSeq2SeqLM.from_pretrained("opus-mt-en-zh")
#使用huggingface/Datasets方式加载数据时,可以用DataCollator达到批处理的效果
data_collator = DataCollatorForSeq2Seq(tokenizer=tokenizer, model=model) #用torch.utils.data.Dataset方式加载时,不需要 training_args = Seq2SeqTrainingArguments(
output_dir="./results",
evaluation_strategy="epoch",
learning_rate=2e-5,
per_device_train_batch_size=16,
per_device_eval_batch_size=16,
weight_decay=0.01,
save_total_limit=3,
num_train_epochs=2,
fp16=True,
)

模型训练

本文以Seq2SeqTrainer作为实例来进行介绍。

trainer = Seq2SeqTrainer(
model=model,
args=training_args,
train_dataset=train_data,
eval_dataset=dev_data,
tokenizer=tokenizer,
data_collator=data_collator, #用torch.utils.data.Dataset方式加载时,此参数不需要
)

补充说明

  • Seq2SeqTrainer()中的train_dataset和eval_dataset参数仅支持torch.utils.data.Dataset和huggingface / Datasets类型的传入实参。

  • torch.utils.data.Dataset类型的实参传入Seq2SeqTrainer()后,会在后序过程直接调用torch.utils.data.DataLoader,与常规pytorch操作相同

  • huggingface / Datasets类型的实参传入Seq2SeqTrainer()后,在后序过程会,先剔除多余的键及其值。至于torch.utils.data.Dataset类型的实参中若包含多余的键及其值,程序会不会报错暂没有测试过。获取模型所需输入参数列表的程序如下:

        def _set_signature_columns_if_needed(self):
    if self._signature_columns is None:
    # Inspect model forward signature to keep only the arguments it accepts.
    signature = inspect.signature(self.model.forward)
    self._signature_columns = list(signature.parameters.keys())
    # Labels may be named label or label_ids, the default data collator handles that.
    self._signature_columns += list(set(["label", "label_ids"] + self.label_names))
  • 即使传递给train_dataset和eval_dataset的数据时字典类型,也存在一种做法使得基于torch.utils.data.Dataset加载数据的方式会报异常

    重载Dataset类(dataset.py)

    # -*- coding: utf-8 -*-
    from torch.utils.data import Dataset class CDNDataset(Dataset):
    def __init__(self, samples):
    super(CDNDataset, self).__init__()
    self.samples = samples def __getitem__(self, ite):
    return self.samples[ite] def __len__(self):
    return len(self.samples)

    main.py中数据读取部分

    train_data = []
    with open('raw_data/txt_en.txt', 'r', encoding='utf-8') as fr_en,
    open('raw_data/txt_zh.txt', 'r', encoding='utf-8') as fr_zh:
    for en_, zh_ in zip(fr_en, fr_zh):
    data = tokenizer(en_.strip(), max_length=128, padding=True, truncation=True, return_tensors='pt')
    data["labels"] = tokenizer(zh_.strip(), max_length=128, padding=True,
    truncation=True, return_tensors='pt')['input_ids']
    train_data.append(data)
    fr_en.close()
    fr_zh.close() train_data = CDNDataset(train_data)
    dev_data = []
    with open('raw_data/test_txt_en.txt', 'r', encoding='utf-8') as fr_en,
    open('raw_data/test_txt_zh.txt', 'r', encoding='utf-8') as fr_zh:
    for en_, zh_ in zip(fr_en, fr_zh):
    data = tokenizer(en_.strip(), max_length=128, padding=True, truncation=True, return_tensors='pt')
    data["labels"] = tokenizer(zh_.strip(), max_length=128, padding=True,
    truncation=True, return_tensors='pt')['input_ids']
    dev_data.append(data)
    fr_en.close()
    fr_zh.close()
    dev_data = CDNDataset(dev_data)

    报错信息:

    "Unable to create tensor, you should probably activate truncation and/or padding "

    ValueError: Unable to create tensor, you should probably activate truncation and/or padding with 'padding=True' 'truncation=True' to have batched tensors with the same length.

    说明:

    现将前一个基于torch.utils.data.Dataset加载数据的方式的案例叫作method1,当前抛出异常的案例叫作method2,两者相比:

    • dataset.py中__getitem__()的返回类型都是字典,每次也都是返回一个样本
    • 在main.py中:
      • method1将所有序列样本存入一个list中,然后对该list进行了一次tokenize,最后在CDNDataset类的__getitem__()中根据索引ite组合成一个样本的字典格式,并返回
      • method2中是先对每个像本序列分别作tokenize,再将各个样本tokenize后得到的字典存入一个list,最后在CDNDataset类的__getitem__()中根据索引ite返回各个样本对应的字典
    • 所报错误信息说没有进行padding 和 truncation, 但事实上我做了,故而不知道是啥问题,望各位大佬不吝赐教。谢过!

transformers 之Trainer对应的数据加载的更多相关文章

  1. ScrollView嵌套ListView,GridView数据加载不全问题的解决

    我们大家都知道ListView,GridView加载数据项,如果数据项过多时,就会显示滚动条.ScrollView组件里面只能包含一个组件,当ScrollView里面嵌套listView,GridVi ...

  2. python多种格式数据加载、处理与存储

    多种格式数据加载.处理与存储 实际的场景中,我们会在不同的地方遇到各种不同的数据格式(比如大家熟悉的csv与txt,比如网页HTML格式,比如XML格式),我们来一起看看python如何和这些格式的数 ...

  3. flask+sqlite3+echarts3+ajax 异步数据加载

    结构: /www | |-- /static |....|-- jquery-3.1.1.js |....|-- echarts.js(echarts3是单文件!!) | |-- /templates ...

  4. Entity Framework关联查询以及数据加载(延迟加载,预加载)

    数据加载分为延迟加载和预加载 EF的关联实体加载有三种方式:Lazy Loading,Eager Loading,Explicit Loading,其中Lazy Loading和Explicit Lo ...

  5. JQuery插件:遮罩+数据加载中。。。(特点:遮你想遮,罩你想罩)

    在很多项目中都会涉及到数据加载.数据加载有时可能会是2-3秒,为了给一个友好的提示,一般都会给一个[数据加载中...]的提示.今天就做了一个这样的提示框. 先去jQuery官网看看怎么写jQuery插 ...

  6. 如何评估ETL的数据加载时间

    简述如何评估大型ETL数据加载时间. 答:评估一个大型的ETL的数据加载时间是一件很复杂的事情.数据加载分为两类,一类是初次加载,另一类是增量加载. 在数据仓库正式投入使用时,需要进行一次初次加载,而 ...

  7. 浅谈Entity Framework中的数据加载方式

    如果你还没有接触过或者根本不了解什么是Entity Framework,那么请看这里http://www.entityframeworktutorial.net/EntityFramework-Arc ...

  8. 实现虚拟模式的动态数据加载Windows窗体DataGridView控件 .net 4.5 (一)

    实现虚拟模式的即时数据加载Windows窗体DataGridView控件 .net 4.5 原文地址 :http://msdn.microsoft.com/en-us/library/ms171624 ...

  9. Android Volley和Gson实现网络数据加载

    Android Volley和Gson实现网络数据加载 先看接口 1 升级接口 http://s.meibeike.com/mcloud/ota/cloudService POST请求 参数列表如下 ...

  10. Echarts通过Ajax实现动态数据加载

    Echarts(3.x版)官网实例的数据都是静态的,实际使用中往往会要求从服务器端取数据进行动态显示,官网教程里给出的异步数据加载很粗略,下面就以官网最简单的实例为例子,详细演示如下过程:1.客户端通 ...

随机推荐

  1. word取消保护

    没有 文档的保护密码 可尝试用此方式,亲测有效 Excel.PPT 应该也可以,没试过 1.新建空白文档 2.插入.对象 3.点击[对象]右边的箭头,选择被加密的文件. 建议两个选项都试一下,我的第二 ...

  2. [Leetcode]移除链表元素

    题目 代码 /** * Definition for singly-linked list. * struct ListNode { * int val; * ListNode *next; * Li ...

  3. 使用SQL4Automation让CodeSYS连接数据库

      摘要:本文旨在说明面向CodeSYS的数据库连接方案SQL4Automation的使用方法. 1.SQL4Automation简介 1.1.什么是SQL4Automation   SQL4Auto ...

  4. Java 进阶P-8.15

    对象串行化 ObjectInputStream类 readObject() ObjectOutputStream类 writeObject() Serializable接口 对象的寿命通常随着生成该对 ...

  5. VS保存后Unity不刷新

    目录 问题:Visual Studio写完代码保存好,Unity不会重新编译 三种解决方案 1.先选为默认.重启Unity.更改为想要的代码编写软件. 2.查看Auto Refresh是否开启 3. ...

  6. 【分析笔记】Linux tasklet 机制的理解

    Tasklet 介绍 Linux 内核提供的四种中断下半部中 softirq(软中断).tasklet(小任务).workqueue(工作队列) .request thread(中断线程)中的其中一种 ...

  7. wsl ubuntu vscode 安装 Fira Code

    如果使用windows terminal(其实就是powershell)那么只需要在windows 中安装 Fira Code 即可,但是如果需要让wsl 中的vscode 也用Fira Code 就 ...

  8. 远程控制 todesk

    最近发现的一个好用的远程连接软件 便是近些年推出来的 todesk 虽然qq的远程 和 向日葵的 远程连接也都可以达到我要实现的效果 但是体验起来的话 我个人还是觉得 todesk更好用一些 下载地址 ...

  9. rt-thread模糊到清晰系列: ipc.c

    #include <rtthread.h> #include <rthw.h> #ifdef RT_USING_HOOK extern void (*rt_object_try ...

  10. 如何获取win10用户最高权限

    第五步,在(输入对象名称)方框中输入"System Managed Accounts Group",再点击"检查名称" 转载: 百度经验:     https: ...