mirror of
https://github.com/modelscope/modelscope.git
synced 2026-08-29 10:08:40 +02:00
[Fix] Fix msdatasets split issue (#1704)
This commit is contained in:
@@ -246,7 +246,7 @@ class MsDataset:
|
||||
'you can trust the external codes.')
|
||||
|
||||
# Raise csv field size limit to avoid errors with large cells
|
||||
if config_kwargs.get('engine') == 'python':
|
||||
if config_kwargs.pop('engine', None) == 'python':
|
||||
import csv as csv_module
|
||||
import sys
|
||||
try:
|
||||
|
||||
@@ -9,12 +9,12 @@ by :func:`~hf_datasets_util.load_dataset_with_ctx`.
|
||||
import importlib
|
||||
import inspect
|
||||
import os
|
||||
import re
|
||||
from functools import partial
|
||||
from pathlib import Path
|
||||
from typing import Dict, List, Optional, Sequence, Tuple, Union
|
||||
from typing import Dict, List, Optional, Tuple
|
||||
|
||||
from datasets import (BuilderConfig, DownloadConfig, DownloadMode, Features,
|
||||
Version, config, data_files)
|
||||
from datasets import (BuilderConfig, DownloadConfig, config)
|
||||
from datasets.data_files import (
|
||||
FILES_TO_IGNORE, DataFilesDict, EmptyDatasetError,
|
||||
_get_data_files_patterns, _is_inside_unrequested_special_dir,
|
||||
@@ -24,17 +24,16 @@ from datasets.download.streaming_download_manager import (
|
||||
_prepare_path_and_storage_options, xbasename, xjoin)
|
||||
from datasets.exceptions import DataFilesNotFoundError
|
||||
from datasets.info import DatasetInfosDict
|
||||
from datasets.load import (BuilderConfigsParameters, DatasetModule,
|
||||
from datasets.load import (BuilderConfigsParameters,
|
||||
DatasetModule,
|
||||
create_builder_configs_from_metadata_configs,
|
||||
get_dataset_builder_class, import_main_class,
|
||||
import_main_class,
|
||||
infer_module_for_data_files)
|
||||
from datasets.naming import camelcase_to_snakecase
|
||||
from datasets.packaged_modules import (_MODULE_TO_EXTENSIONS,
|
||||
_PACKAGED_DATASETS_MODULES)
|
||||
from datasets.utils.file_utils import (cached_path, is_local_path,
|
||||
relative_to_absolute_path)
|
||||
from datasets.utils.file_utils import (cached_path, is_local_path)
|
||||
from datasets.utils.metadata import MetadataConfigs
|
||||
from datasets.utils.track import tracked_str
|
||||
from fsspec import filesystem
|
||||
from fsspec.core import _un_chain
|
||||
from fsspec.utils import stringify_path
|
||||
@@ -42,14 +41,16 @@ from huggingface_hub import DatasetCard, DatasetCardData
|
||||
from packaging import version
|
||||
|
||||
from modelscope import HubApi
|
||||
from modelscope.msdatasets.utils._compat import (
|
||||
_HAS_SCRIPT_LOADING, _create_importable_file, _get_importable_file_path,
|
||||
_load_importable_file, files_to_hash, get_imports, init_dynamic_modules,
|
||||
from modelscope.msdatasets.utils._compat import (_create_importable_file,
|
||||
_get_importable_file_path,
|
||||
_load_importable_file,
|
||||
files_to_hash,
|
||||
get_imports,
|
||||
init_dynamic_modules,
|
||||
resolve_trust_remote_code)
|
||||
from modelscope.utils.constant import (DEFAULT_DATASET_REVISION,
|
||||
REPO_TYPE_DATASET)
|
||||
from modelscope.utils.file_utils import is_relative_path
|
||||
from modelscope.utils.import_utils import has_attr_in_class
|
||||
from modelscope.utils.logger import get_logger
|
||||
|
||||
# ALL_ALLOWED_EXTENSIONS moved to datasets.packaged_modules in datasets 4.0
|
||||
@@ -60,6 +61,73 @@ except ImportError:
|
||||
|
||||
logger = get_logger()
|
||||
|
||||
|
||||
def _extract_split_names(split):
|
||||
"""Extract base split names from a split specification string.
|
||||
|
||||
Handles simple names ("tool"), sliced splits ("train[:100]"),
|
||||
and combined splits ("train+test").
|
||||
|
||||
Args:
|
||||
split: A split specification string, or None.
|
||||
|
||||
Returns:
|
||||
A set of split name strings, or None if *split* is None or
|
||||
cannot be parsed.
|
||||
"""
|
||||
if split is None:
|
||||
return None
|
||||
split_str = str(split)
|
||||
parts = split_str.split('+')
|
||||
names = set()
|
||||
for part in parts:
|
||||
# Remove slice notation like "[:100]" or "[50%:]"
|
||||
name = re.split(r'\[', part.strip())[0]
|
||||
if name:
|
||||
names.add(name)
|
||||
return names if names else None
|
||||
|
||||
|
||||
def _filter_data_files_by_split(data_files, download_config):
|
||||
"""Filter data_files entries to only include the requested split(s).
|
||||
|
||||
Args:
|
||||
data_files: The raw data_files value from metadata_configs.
|
||||
Expected format: list of dicts with 'split' and 'path' keys.
|
||||
download_config: The DownloadConfig instance (may be None).
|
||||
|
||||
Returns:
|
||||
Filtered data_files if a split filter is active; otherwise
|
||||
the original *data_files* unchanged.
|
||||
"""
|
||||
# 1. Safely retrieve the split value
|
||||
split_str = None
|
||||
if download_config is not None:
|
||||
storage_opts = getattr(download_config, 'storage_options', None)
|
||||
if isinstance(storage_opts, dict):
|
||||
split_str = storage_opts.get('split')
|
||||
|
||||
if split_str is None:
|
||||
return data_files
|
||||
|
||||
# 2. Parse split names
|
||||
split_names = _extract_split_names(split_str)
|
||||
if not split_names:
|
||||
return data_files
|
||||
|
||||
# 3. Only filter list[dict] format
|
||||
if not isinstance(data_files, list):
|
||||
return data_files
|
||||
if not all(isinstance(item, dict) and 'split' in item for item in data_files):
|
||||
return data_files
|
||||
|
||||
# 4. Filter
|
||||
filtered = [df for df in data_files if df.get('split') in split_names]
|
||||
|
||||
# 5. Fallback: if filtered is empty, return original data
|
||||
return filtered if filtered else data_files
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Shared HubApi instance (avoids creating a new requests.Session per call)
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -575,6 +643,8 @@ def get_module_without_script(self) -> DatasetModule:
|
||||
else:
|
||||
subset_data_files = next(iter(
|
||||
metadata_configs.values()))['data_files']
|
||||
subset_data_files = _filter_data_files_by_split(
|
||||
subset_data_files, self.download_config)
|
||||
patterns = sanitize_patterns(subset_data_files)
|
||||
else:
|
||||
patterns = _get_data_patterns(
|
||||
|
||||
@@ -25,7 +25,7 @@ from datasets import (Dataset, DatasetBuilder, DatasetDict,
|
||||
DownloadConfig, DownloadManager, DownloadMode, Features,
|
||||
IterableDataset, IterableDatasetDict, Split,
|
||||
VerificationMode, Version, config, data_files, LargeList,
|
||||
Sequence as SequenceHf)
|
||||
Sequence as SequenceHf, SplitDict)
|
||||
|
||||
try:
|
||||
from datasets import List as DatasetList
|
||||
@@ -456,6 +456,81 @@ def _hf_fs_open(self, path, mode='rb', **kwargs):
|
||||
return _hf_fs_open_original(self, path, mode=mode, **kwargs)
|
||||
|
||||
|
||||
def _validate_split_exists(builder_instance, split):
|
||||
"""Fail-fast check: raise ValueError before downloading if the
|
||||
requested split does not exist in the dataset metadata.
|
||||
|
||||
Args:
|
||||
builder_instance: The DatasetBuilder instance with info/config.
|
||||
split: The user-requested split specification (may be None).
|
||||
|
||||
Raises:
|
||||
ValueError: If any requested split name is not found among
|
||||
the available splits declared in the dataset metadata.
|
||||
"""
|
||||
if split is None:
|
||||
return
|
||||
|
||||
from modelscope.msdatasets.utils._module_factories import _extract_split_names
|
||||
split_names = _extract_split_names(split)
|
||||
if not split_names:
|
||||
return
|
||||
|
||||
# Prefer info.splits (original metadata); fall back to data_files keys
|
||||
available = set()
|
||||
info = getattr(builder_instance, 'info', None)
|
||||
if info is not None and info.splits:
|
||||
available = set(info.splits.keys())
|
||||
|
||||
if not available:
|
||||
config = getattr(builder_instance, 'config', None)
|
||||
data_files = getattr(config, 'data_files', None)
|
||||
if isinstance(data_files, dict):
|
||||
available = set(data_files.keys())
|
||||
|
||||
if not available:
|
||||
return # Cannot determine available splits; let downstream handle
|
||||
|
||||
missing = split_names - available
|
||||
if missing:
|
||||
raise ValueError(
|
||||
f'Split {sorted(missing)} not found in dataset. '
|
||||
f'Available splits: {sorted(available)}'
|
||||
)
|
||||
|
||||
|
||||
def _align_builder_splits_with_data_files(builder_instance, split):
|
||||
"""Align builder.info.splits with the actually requested split(s).
|
||||
|
||||
When data_files have been filtered to a subset of splits (see
|
||||
_filter_data_files_by_split in _module_factories.py), the builder's
|
||||
info.splits metadata may still list all original splits from the
|
||||
README. download_and_prepare() calls verify_splits() which would
|
||||
then raise ExpectedMoreSplitsError. This helper prunes info.splits
|
||||
to only contain the splits that will actually be generated.
|
||||
"""
|
||||
if split is None:
|
||||
return
|
||||
info = getattr(builder_instance, 'info', None)
|
||||
if info is None or info.splits is None:
|
||||
return
|
||||
|
||||
from modelscope.msdatasets.utils._module_factories import _extract_split_names
|
||||
split_names = _extract_split_names(split)
|
||||
if not split_names:
|
||||
return
|
||||
|
||||
existing_keys = set(info.splits.keys())
|
||||
if split_names >= existing_keys:
|
||||
return # All splits requested, no filtering needed
|
||||
|
||||
filtered = {k: v for k, v in info.splits.items() if k in split_names}
|
||||
if not filtered:
|
||||
return # Safety: don't empty out splits
|
||||
|
||||
info.splits = SplitDict(filtered, dataset_name=info.splits.dataset_name)
|
||||
|
||||
|
||||
# ===================================================================
|
||||
# DatasetsWrapperHF
|
||||
# ===================================================================
|
||||
@@ -544,6 +619,7 @@ class DatasetsWrapperHF:
|
||||
storage_options=storage_options,
|
||||
trust_remote_code=trust_remote_code,
|
||||
_require_default_config_name=name is None,
|
||||
split=split,
|
||||
**config_kwargs,
|
||||
)
|
||||
|
||||
@@ -572,6 +648,9 @@ class DatasetsWrapperHF:
|
||||
if streaming:
|
||||
return builder_instance.as_streaming_dataset(split=split)
|
||||
|
||||
_validate_split_exists(builder_instance, split)
|
||||
_align_builder_splits_with_data_files(builder_instance, split)
|
||||
|
||||
builder_instance.download_and_prepare(
|
||||
download_config=download_config,
|
||||
download_mode=download_mode,
|
||||
@@ -624,6 +703,7 @@ class DatasetsWrapperHF:
|
||||
storage_options: Optional[Dict] = None,
|
||||
trust_remote_code: Optional[bool] = None,
|
||||
_require_default_config_name=True,
|
||||
split: Optional[Union[str, Split]] = None,
|
||||
**config_kwargs,
|
||||
) -> DatasetBuilder:
|
||||
|
||||
@@ -644,6 +724,12 @@ class DatasetsWrapperHF:
|
||||
download_config = download_config.copy(
|
||||
) if download_config else DownloadConfig()
|
||||
download_config.storage_options.update(storage_options)
|
||||
if split is not None:
|
||||
download_config = download_config.copy(
|
||||
) if download_config else DownloadConfig()
|
||||
if download_config.storage_options is None:
|
||||
download_config.storage_options = {}
|
||||
download_config.storage_options['split'] = split
|
||||
|
||||
dataset_module = DatasetsWrapperHF.dataset_module_factory(
|
||||
path,
|
||||
|
||||
@@ -53,11 +53,12 @@ def _request_with_retry_ms(
|
||||
url: str,
|
||||
max_retries: int = 2,
|
||||
base_wait_time: float = 0.5,
|
||||
max_wait_time: float = 2,
|
||||
max_wait_time: float = 8,
|
||||
timeout: float = 10.0,
|
||||
**params,
|
||||
) -> requests.Response:
|
||||
"""Wrapper around requests to retry in case it fails with a ConnectTimeout, with exponential backoff.
|
||||
"""Wrapper around requests to retry in case it fails with a ConnectTimeout,
|
||||
ReadTimeout or ConnectionError, with exponential backoff.
|
||||
|
||||
Note that if the environment variable HF_DATASETS_OFFLINE is set to 1, then a OfflineModeIsEnabled error is raised.
|
||||
|
||||
@@ -77,7 +78,9 @@ def _request_with_retry_ms(
|
||||
try:
|
||||
response = requests.request(method=method.upper(), url=url, timeout=timeout, **params)
|
||||
success = True
|
||||
except (requests.exceptions.ConnectTimeout, requests.exceptions.ConnectionError) as err:
|
||||
except (requests.exceptions.ConnectTimeout,
|
||||
requests.exceptions.ConnectionError,
|
||||
requests.exceptions.ReadTimeout) as err:
|
||||
if tries > max_retries:
|
||||
raise err
|
||||
else:
|
||||
@@ -88,7 +91,7 @@ def _request_with_retry_ms(
|
||||
|
||||
|
||||
def http_head_ms(
|
||||
url, proxies=None, headers=None, cookies=None, allow_redirects=True, timeout=10.0, max_retries=0
|
||||
url, proxies=None, headers=None, cookies=None, allow_redirects=True, timeout=10.0, max_retries=3
|
||||
) -> requests.Response:
|
||||
headers = copy.deepcopy(headers) or {}
|
||||
headers['user-agent'] = get_datasets_user_agent_ms(user_agent=headers.get('user-agent'))
|
||||
@@ -106,7 +109,7 @@ def http_head_ms(
|
||||
|
||||
|
||||
def http_get_ms(
|
||||
url, temp_file, proxies=None, resume_size=0, headers=None, cookies=None, timeout=100.0, max_retries=0, desc=None
|
||||
url, temp_file, proxies=None, resume_size=0, headers=None, cookies=None, timeout=300.0, max_retries=3, desc=None
|
||||
) -> Optional[requests.Response]:
|
||||
headers = dict(headers) if headers is not None else {}
|
||||
headers['user-agent'] = get_datasets_user_agent_ms(user_agent=headers.get('user-agent'))
|
||||
@@ -147,7 +150,7 @@ def get_from_cache_ms(
|
||||
user_agent=None,
|
||||
local_files_only=False,
|
||||
use_etag=True,
|
||||
max_retries=0,
|
||||
max_retries=3,
|
||||
token=None,
|
||||
use_auth_token='deprecated',
|
||||
ignore_url_params=False,
|
||||
|
||||
Reference in New Issue
Block a user