Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
37 changes: 1 addition & 36 deletions dataflow/utils/storage.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
# 读出当前批次数据
Expand Down
133 changes: 133 additions & 0 deletions test/cpu_only/test_remote_batch_slicing.py
Original file line number Diff line number Diff line change
@@ -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()
Loading