diff --git a/dataflow/utils/storage.py b/dataflow/utils/storage.py index ac0a82a9..9e30f6f9 100644 --- a/dataflow/utils/storage.py +++ b/dataflow/utils/storage.py @@ -973,45 +973,10 @@ def read(self, output_type: Literal["dataframe", "dict"]="dataframe") -> Any: file_path = self._get_cache_file_path(self.operator_step) self.logger.info(f"Reading data from {file_path} with type {output_type}") - if self.operator_step == 0: - source = self.first_entry_file_name - self.logger.info(f"Reading remote dataset from {source} with type {output_type}") - if source.startswith("hf:"): - from datasets import load_dataset - _, dataset_name, *parts = source.split(":") - - if len(parts) == 1: - config, split = None, parts[0] - elif len(parts) == 2: - config, split = parts - else: - config, split = None, "train" - - dataset = ( - load_dataset(dataset_name, config, split=split) - if config - else load_dataset(dataset_name, split=split) - ) - dataframe = dataset.to_pandas() - return self._convert_output(dataframe, output_type) - - elif source.startswith("ms:"): - from modelscope import MsDataset - _, dataset_name, *split_parts = source.split(":") - split = split_parts[0] if split_parts else "train" - - dataset = MsDataset.load(dataset_name, split=split) - dataframe = pd.DataFrame(dataset) - return self._convert_output(dataframe, output_type) - - else: - local_cache = file_path.split(".")[-1] - else: - local_cache = self.cache_type if self._dataframe_buffer.get(self.operator_step) is not None: dataframe = self._dataframe_buffer[self.operator_step].copy() else: - dataframe = self._load_local_file(file_path, local_cache) + dataframe = super().read(output_type="dataframe") self._dataframe_buffer[self.operator_step] = dataframe.copy() self.record_count = len(dataframe) # 读出当前批次数据 diff --git a/test/cpu_only/test_remote_batch_slicing.py b/test/cpu_only/test_remote_batch_slicing.py new file mode 100644 index 00000000..c592cb45 --- /dev/null +++ b/test/cpu_only/test_remote_batch_slicing.py @@ -0,0 +1,133 @@ +import sys +import types +from unittest.mock import Mock + +import pandas as pd +import pytest + +from dataflow.utils.storage import BatchedFileStorage, StreamBatchedFileStorage + + +def mock_remote_source(monkeypatch, source_kind, dataframe): + if source_kind == "hf": + loader = Mock(return_value=types.SimpleNamespace(to_pandas=dataframe.copy)) + module = types.ModuleType("datasets") + module.load_dataset = loader + monkeypatch.setitem(sys.modules, "datasets", module) + source = "hf:example/data:subset:validation" + else: + loader = Mock(return_value=dataframe.to_dict(orient="records")) + module = types.ModuleType("modelscope") + module.MsDataset = types.SimpleNamespace(load=loader) + monkeypatch.setitem(sys.modules, "modelscope", module) + source = "ms:example/data:validation" + return source, loader + + +@pytest.mark.cpu +@pytest.mark.parametrize("source_kind", ["hf", "ms"]) +@pytest.mark.parametrize("output_type", ["dataframe", "dict"]) +def test_remote_batches_are_sliced_counted_and_cached( + monkeypatch, tmp_path, source_kind, output_type +): + source, loader = mock_remote_source( + monkeypatch, + source_kind, + pd.DataFrame({"text": ["first", "second", "third"]}, index=[10, 20, 30]), + ) + storage = BatchedFileStorage(source, cache_path=str(tmp_path)).step() + storage.batch_size = 2 + + for batch_step, expected in enumerate([["first", "second"], ["third"], []]): + storage.batch_step = batch_step + result = storage.read(output_type) + if output_type == "dataframe": + assert result["text"].tolist() == expected + assert result.index.tolist() == list(range(len(expected))) + if not result.empty: + result.loc[0, "text"] = "modified by caller" + else: + assert result == [{"text": text} for text in expected] + if result: + result[0]["text"] = "modified by caller" + assert storage.record_count == 3 + + storage.batch_step = 0 + assert storage.read("dict") == [{"text": "first"}, {"text": "second"}] + if source_kind == "hf": + loader.assert_called_once_with("example/data", "subset", split="validation") + else: + loader.assert_called_once_with("example/data", split="validation") + + +@pytest.mark.cpu +@pytest.mark.parametrize("source_kind", ["hf", "ms"]) +def test_remote_read_without_batch_size_keeps_all_rows(monkeypatch, tmp_path, source_kind): + source, loader = mock_remote_source( + monkeypatch, source_kind, pd.DataFrame({"text": ["first", "second", "third"]}) + ) + storage = BatchedFileStorage(source, cache_path=str(tmp_path)).step() + + assert storage.read("dict") == [ + {"text": "first"}, {"text": "second"}, {"text": "third"} + ] + assert storage.record_count == 3 + loader.assert_called_once() + + +@pytest.mark.cpu +@pytest.mark.parametrize("source_kind", ["hf", "ms"]) +def test_empty_remote_dataset_has_zero_records(monkeypatch, tmp_path, source_kind): + source, loader = mock_remote_source(monkeypatch, source_kind, pd.DataFrame()) + storage = BatchedFileStorage(source, cache_path=str(tmp_path)).step() + storage.batch_size = 2 + + assert storage.read().empty + assert storage.read("dict") == [] + assert storage.record_count == 0 + loader.assert_called_once() + + +@pytest.mark.cpu +@pytest.mark.parametrize("storage_class", [BatchedFileStorage, StreamBatchedFileStorage]) +@pytest.mark.parametrize("extension", ["jsonl", "csv"]) +def test_local_batch_reads_keep_slicing_and_caching(tmp_path, storage_class, extension): + source = tmp_path / f"source.{extension}" + dataframe = pd.DataFrame({"text": ["first", "second", "third"]}) + if extension == "jsonl": + dataframe.to_json(source, orient="records", lines=True) + else: + dataframe.to_csv(source, index=False) + storage = storage_class(str(source), cache_path=str(tmp_path)).step() + storage.batch_size = 2 + storage.batch_step = 1 + + result = storage.read() + assert result["text"].tolist() == ["third"] + assert result.index.tolist() == [0] + assert storage.record_count == 3 + + # Cached reads must not reload a source that has already been read. + source.unlink() + storage.batch_step = 0 + assert storage.read("dict") == [{"text": "first"}, {"text": "second"}] + + +@pytest.mark.cpu +@pytest.mark.parametrize("source_kind", ["hf", "ms"]) +def test_later_step_reads_local_output_without_reloading_remote_source( + monkeypatch, tmp_path, source_kind +): + source, loader = mock_remote_source( + monkeypatch, source_kind, pd.DataFrame({"text": ["first", "second", "third"]}) + ) + storage = BatchedFileStorage(source, cache_path=str(tmp_path)).step() + storage.read() + storage.write(pd.DataFrame({"text": ["processed first", "processed second", "processed third"]})) + next_step = storage.step() + next_step.batch_size = 2 + next_step.batch_step = 1 + + assert next_step.read("dict") == [{"text": "processed third"}] + assert next_step.record_count == 3 + loader.assert_called_once()