You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何无需torchdata的IterableWrapper封装使用Huggingface Trainer处理流式数据集?

问题:流式IterableDataset直接用于Seq2SeqTrainer触发报错

场景复现

通过stream=True加载流式数据集:

train_data = load_dataset("csv", data_files="../input/tatoeba/tatoeba-sentpairs.tsv", 
                  streaming=True, delimiter="\t", split="train")

尝试直接将该IterableDataset传入Seq2SeqTrainer:

# 初始化Trainer
trainer = Seq2SeqTrainer(
    model=multibert,
    tokenizer=tokenizer,
    args=training_args,
    train_dataset=train_data,
    eval_dataset=train_data,
)

trainer.train()

报错信息

运行后触发如下类型错误:

---------------------------------------------------------------------------
TypeError                                 Traceback (most recent call last)
/tmp/ipykernel_27/3002801805.py in <module>
     28 )
     29 
---> 30 trainer.train()

/opt/conda/lib/python3.7/site-packages/transformers/trainer.py in train(self, resume_from_checkpoint, trial, ignore_keys_for_eval, **kwargs)
   1411             resume_from_checkpoint=resume_from_checkpoint,
   1412             trial=trial,
-> 1413             ignore_keys_for_eval=ignore_keys_for_eval,
   1414         )
   1415 

/opt/conda/lib/python3.7/site-packages/transformers/trainer.py in _inner_training_loop(self, batch_size, args, resume_from_checkpoint, trial, ignore_keys_for_eval)
   1623 
   1624             step = -1
-> 1625             for step, inputs in enumerate(epoch_iterator):
   1626 
   1627                 # Skip past any already trained steps if resuming training

/opt/conda/lib/python3.7/site-packages/torch/utils/data/dataloader.py in __next__(self)
    528             if self._sampler_iter is None:
    529                 self._reset()
-> 530             data = self._next_data()
    531             self._num_yielded += 1
    532             if self._dataset_kind == _DatasetKind.Iterable and \

/opt/conda/lib/python3.7/site-packages/torch/utils/data/dataloader.py in _next_data(self)
    567 
    568     def _next_data(self):
-> 569         index = self._next_index()  # may raise StopIteration
    570         data = self._dataset_fetcher.fetch(index)  # may raise StopIteration
    571         if self._pin_memory:

/opt/conda/lib/python3.7/site-packages/torch/utils/data/dataloader.py in _next_index(self)
    519 
    520     def _next_index(self):
-> 521         return next(self._sampler_iter)  # may raise StopIteration
    522 
    523     def _next_data(self):

/opt/conda/lib/python3.7/site-packages/torch/utils/data/sampler.py in __iter__(self)
    224     def __iter__(self) -> Iterator[List[int]]:
    225         batch = []
-> 226         for idx in self.sampler:
    227             batch.append(idx)
    228             if len(batch) == self.batch_size:

/opt/conda/lib/python3.7/site-packages/torch/utils/data/sampler.py in __iter__(self)
     64 
     65     def __iter__(self) -> Iterator[int]:
---&gt; 66         return iter(range(len(self.data_source)))
     67 
     68     def __len__(self) -> int:

TypeError: object of type 'IterableDataset' has no len()

现有解决方案:用IterableWrapper封装

通过torchdata库的IterableWrapper包装流式数据集,可解决该报错:

from torchdata.datapipes.iter import IterableWrapper

...

# 初始化Trainer
trainer = Seq2SeqTrainer(
    model=multibert,
    tokenizer=tokenizer,
    args=training_args,
    train_dataset=IterableWrapper(train_data),
    eval_dataset=IterableWrapper(train_data),
)

trainer.train()

核心疑问

能否不通过IterableWrapper转换,直接将IterableDataset与Seq2SeqTrainer配合使用?


完整可复现代码

将代码中train_dataset=IterableWrapper(train_data)替换为train_dataset=train_data,即可复现上述报错:

import torch

from datasets import load_dataset
from transformers import EncoderDecoderModel
from transformers import AutoTokenizer
from transformers import Seq2SeqTrainer, Seq2SeqTrainingArguments

from torchdata.datapipes.iter import IterableWrapper

multibert = EncoderDecoderModel.from_encoder_decoder_pretrained(
    "bert-base-multilingual-uncased", "bert-base-multilingual-uncased"
)
tokenizer= AutoTokenizer.from_pretrained("bert-base-multilingual-uncased")
tokenizer.bos_token = tokenizer.cls_token
tokenizer.eos_token = tokenizer.sep_token
tokenizer.add_special_tokens({'pad_token': '[PAD]'})

# 设置特殊Token
multibert.config.decoder_start_token_id = tokenizer.bos_token_id
multibert.config.eos_token_id = tokenizer.eos_token_id
multibert.config.pad_token_id = tokenizer.pad_token_id

# 设置beam搜索参数
multibert.config.vocab_size = multibert.config.decoder.vocab_size

def process_data_to_model_inputs(batch, max_len=10): 
    inputs = tokenizer(batch["SRC"], padding="max_length",
                       truncation=True, max_length=max_len)
    outputs = tokenizer(batch["TRG"], padding="max_length", 
                        truncation=True, max_length=max_len)

    batch["input_ids"] = inputs.input_ids
    batch["attention_mask"] = inputs.attention_mask
    batch["decoder_input_ids"] = outputs.input_ids
    batch["decoder_attention_mask"] = outputs.attention_mask
    batch["labels"] = outputs.input_ids.copy()

    # BERT会自动偏移标签,因此labels与decoder_input_ids完全对应
    # 需要确保PAD Token被忽略
    batch["labels"] = [[-100 if token == tokenizer.pad_token_id else token for token in labels] for labels in batch["labels"]]
    
    return batch


# tatoeba-sentpairs.tsv是一个较大的文件
train_data = load_dataset("csv", data_files="../input/tatoeba/tatoeba-sentpairs.tsv", 
                  streaming=True, delimiter="\t", split="train")

train_data = train_data.map(process_data_to_model_inputs, batched=True)


batch_size = 1

# 设置训练参数(未调优,可自行修改)
training_args = Seq2SeqTrainingArguments(
    output_dir="./",
    evaluation_strategy="steps",
    per_device_train_batch_size=batch_size,
    per_device_eval_batch_size=batch_size,
    predict_with_generate=True,
    logging_steps=2,  # 全量训练时设为1000
    save_steps=16,    # 全量训练时设为500
    eval_steps=4,     # 全量训练时设为8000
    warmup_steps=1,   # 全量训练时设为2000
    max_steps=16,     # 全量训练时删除该行
    # overwrite_output_dir=True,
    save_total_limit=1,
    #fp16=True, 
)


# 初始化Trainer
trainer = Seq2SeqTrainer(
    model=multibert,
    tokenizer=tokenizer,
    args=training_args,
    train_dataset=IterableWrapper(train_data),
    eval_dataset=IterableWrapper(train_data),
)

trainer.train()

内容的提问来源于stack exchange,提问作者alvas

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.14 23:55:37