如何在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
相关产品推荐
相关产品推荐

