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

如何在Hugging Face非流式Dataset中使用take()方法?

非流式HF Dataset无法调用.take()方法的解决办法

核心报错

仅在streaming=False加载Hugging Face数据集时,触发如下错误:

Exception has occurred: AttributeError       (note: full exception trace is shown but execution is paused at: _run_module_as_main)
'Dataset' object has no attribute 'take'
  File "/lfs/ampere1/0/brando9/beyond-scale-language-data-diversity/src/diversity/div_coeff.py", line 499, in experiment_compute_diveristy_coeff_single_dataset_then_combined_datasets_with_domain_weights
    batch = dataset.take(batch_size)
  File "/lfs/ampere1/0/brando9/beyond-scale-language-data-diversity/src/diversity/div_coeff.py", line 552, in <module>
    experiment_compute_diveristy_coeff_single_dataset_then_combined_datasets_with_domain_weights()
  File "/lfs/ampere1/0/brando9/miniconda/envs/beyond_scale/lib/python3.10/runpy.py", line 86, in _run_code
    exec(code, run_globals)
  File "/lfs/ampere1/0/brando9/miniconda/envs/beyond_scale/lib/python3.10/runpy.py", line 196, in _run_module_as_main (Current frame)
    return _run_code(code, main_globals, None,
AttributeError: 'Dataset' object has no attribute 'take'

原因:非流式加载的arrow_dataset.Dataset没有.take()方法,该方法仅属于流式的IterableDataset。

已尝试的方案

  • 方案1:转换为IterableDataset
    通过强制转换类型来调用.take(),代码如下,但存在获取批次速度极慢的问题:
    print(f'{dataset=}')
    print(f'{type(dataset)=}')
    # datasets.iterable_dataset.IterableDataset
    # datasets.arrow_dataset.Dataset
    dataset = IterableDataset(dataset) if type(dataset) != IterableDataset else dataset  # to force dataset.take(batch_size) to work in non-streaming mode
    batch = dataset.take(batch_size)
    
  • 方案2:自定义Collate函数
    尝试通过自定义Collate函数直接获取指定大小批次,但未成功,且不想使用Hugging Face Trainer。

额外问题:流式请求连接中断

使用流式数据时还会触发连接中断错误:

2654     httplib_response = self._make_request(
2655   File "/lfs/hyperturing1/0/brando9/miniconda/envs/beyond_scale/lib/python3.10/site-packages/urllib3/connectionpool.py", line 466, in _make_request
2656     six.raise_from(e, None)
2657   File "<string>", line 3, in raise_from
2658   File "/lfs/hyperturing1/0/brando9/miniconda/envs/beyond_scale/lib/python3.10/site-packages/urllib3/connectionpool.py", line 461, in _make_request
2659     httplib_response = conn.getresponse()
2660   File "/lfs/hyperturing1/0/brando9/miniconda/envs/beyond_scale/lib/python3.10/http/client.py", line 1375, in getresponse
2661     response.begin()
2662   File "/lfs/hyperturing1/0/brando9/miniconda/envs/beyond_scale/lib/python3.10/http/client.py", line 318, in begin
2663     version, status, reason = self._read_status()
2664   File "/lfs/hyperturing1/0/brando9/miniconda/envs/beyond_scale/lib/python3.10/http/client.py", line 287, in _read_status
2665     raise RemoteDisconnected("Remote end closed connection without"
2666 http.client.RemoteDisconnected: Remote end closed connection without response
2667 During handling of the above exception, another exception occurred:
2668 Traceback (most recent call last):
2669   File "/lfs/hyperturing1/0/brando9/miniconda/envs/beyond_scale/lib/python3.10/site-packages/requests/adapters.py", line 486, in send
2670     resp = conn.urlopen(
2671   File "/lfs/hyperturing1/0/brando9/miniconda/envs/beyond_scale/lib/python3.10/site-packages/urllib3/connectionpool.py", line 798, in urlopen
2672     retries = retries.increment(
2673   File "/lfs/hyperturing1/0/brando9/miniconda/envs/beyond_scale/lib/python3.10/site-packages/urllib3/util/retry.py", line 550, in increment
2674     raise six.reraise(type(error), error, _stacktrace)
2675   File "/lfs/hyperturing1/0/brando9/miniconda/envs/beyond_scale/lib/python3.10/site-packages/urllib3/packages/six.py", line 769, in reraise
2676     raise value.with_traceback(tb)
2677   File "/lfs/hyperturing1/0/brando9/miniconda/envs/beyond_scale/lib/python3.10/site-packages/urllib3/connectionpool.py", line 714, in urlopen
2678     httplib_response = self._make_request(
2679   File "/lfs/hyperturing1/0/brando9/miniconda/envs/beyond_scale/lib/python3.10/site-packages/urllib3/connectionpool.py", line 466, in _make_request
2680     six.raise_from(e, None)
2681   File "<string>", line 3, in raise_from
2682   File "/lfs/hyperturing1/0/brando9/miniconda/envs/beyond_scale/lib/python3.10/site-packages/urllib3/connectionpool.py", line 461, in _make_request
2683     httplib_response = conn.getresponse()
2684   File "/lfs/hyperturing1/0/brando9/miniconda/envs/beyond_scale/lib/python3.10/http/client.py", line 1375, in getresponse
2685     response.begin()
2686   File "/lfs/hyperturing1/0/brando9/miniconda/envs/beyond_scale/lib/python3.10/http/client.py", line 318, in begin
2687     version, status, reason = self._read_status()
2688   File "/lfs/hyperturing1/0/brando9/miniconda/envs/beyond_scale/lib/python3.10/http/client.py", line 287, in _read_status
2689     raise RemoteDisconnected("Remote end closed connection without"
2690 urllib3.exceptions.ProtocolError: ('Connection aborted.', RemoteDisconnected('Remote end closed connection without response'))
2691 During handling of the above exception, another exception occurred:
2692 Traceback (most recent call last):
2693   File "/lfs/hyperturing1/0/brando9/beyond-scale-language-data-diversity/src/diversity/div_coeff.py", line 578, in <module>
2694     # -- Finish wandb
2695   File "/lfs/hyperturing1/0/brando9/beyond-scale-language-data-diversity/src/diversity/div_coeff.py", line 540, in experiment_compute_diveristy_coeff_single_dataset_then_combined_datasets_with_domain_weights
2696     print(f'{batch=}')
2697   File "/lfs/hyperturing1/0/brando9/beyond-scale-language-data-diversity/src/diversity/div_coeff.py", line 63, in get_diversity_coefficient
2698     embedding, loss = Task2Vec(probe_network, classifier_opts={'seed': seed}).embed(tokenized_batch)
2699   File "/afs/cs.stanford.edu/u/brando9/beyond-scale-language-data-diversity/src/diversity/task2vec.py", line 133, in embed
2700     loss = self._finetune_classifier(dataset, loader_opts=self.loader_opts, classifier_opts=self.classifier_opts, max_samples=self.max_samples, epochs=epochs)
2701   File "/afs/cs.stanford.edu/u/brando9/beyond-scale-language-data-diversity/src/diversity/task2vec.py", line 198, in _finetune_classifier
2702     for step, batch in enumerate(epoch_iterator):
2703   File "/lfs/hyperturing1/0/brando9/miniconda/envs/beyond_scale/lib/python3.10/site-packages/tqdm/std.py", line 1182, in __iter__
2704     for obj in iterable:
2705   File "/lfs/hyperturing1/0/brando9/miniconda/envs/beyond_scale/lib/python3.10/site-packages/torch/utils/data/dataloader.py", line 633, in __next__
2706     data = self._next_data()
2707   File "/lfs/hyperturing1/0/brando9/miniconda/envs/beyond_scale/lib/python3.10/site-packages/torch/utils/data/dataloader.py", line 677, in _next_data
2708     data = self._dataset_fetcher.fetch(index)  # may raise StopIteration
2709   File "/lfs/hyperturing1/0/brando9/miniconda/envs/beyond_scale/lib/python3.10/site-packages/torch/utils/data/_utils/fetch.py", line 32, in fetch
2710     data.append(next(self.dataset_iter))
2711   File "/lfs/hyperturing1/0/brando9/miniconda/envs/beyond_scale/lib/python3.10/site-packages/datasets/iterable_dataset.py", line 1353, in __iter__
2712     for key, example in ex_iterable:
2713   File "/lfs/hyperturing1/0/brando9/miniconda/envs/beyond_scale/lib/python3.10/site-packages/datasets/iterable_dataset.py", line 652, in __iter__
2714     yield from self._iter()
2715   File "/lfs/hyperturing1/0/brando9/miniconda/envs/beyond_scale/lib/python3.10/site-packages/datasets/iterable_dataset.py", line 667, in _iter
2716     for key, example in iterator:

有效解决方案

1. 替代.take()的快速方法

非流式的arrow_dataset.Dataset支持直接切片或select()方法,完全可以替代.take(),且速度远快于转换为IterableDataset:

  • 切片方式:
    batch = dataset[:batch_size]
    
  • select方式(适合需要指定非连续索引的场景):
    batch = dataset.select(range(batch_size))
    

2. 解决流式请求中断问题

  • 加载数据集时增加重试次数:
    from datasets import load_dataset, DownloadConfig
    
    download_config = DownloadConfig(max_retries=5)
    dataset = load_dataset("your_dataset_name", streaming=True, download_config=download_config)
    
  • 提前将数据集下载到本地,避免流式请求:
    # 先下载到本地缓存
    dataset = load_dataset("your_dataset_name", streaming=False)
    # 后续直接从本地加载
    dataset = load_dataset("your_dataset_name", cache_dir="/path/to/cache")
    
  • 设置环境变量延长缓存超时:
    export HF_DATASETS_CACHE_TIMEOUT=3600
    

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.13 14:55:54