diff --git a/README.md b/README.md index 3c08a6ad..8482d9f2 100644 --- a/README.md +++ b/README.md @@ -116,9 +116,17 @@ To simplify the usage of TransferQueue, we provide a Redis-style high-level API - **(async_)kv_put**: Insert/Update a multi-column sample by key, with an optional metadata tag. - **(async_)kv_batch_put**: Put multiple key-value pairs efficiently in batches. - **(async_)kv_batch_get**: Retrieve samples (by keys), supporting column selection (by fields). +- **(async_)kv_update / (async_)kv_batch_update**: Merge new values into already-produced fields with `merge_fn(old, new)`. SimpleStorage only. +- **(async_)kv_empty**: Store `None` in selected fields while keeping their keys ready. SimpleStorage only. - **(async_)kv_list**: List keys and tags (metadata) in a partition. - **(async_)kv_clear**: Remove key-value pairs from storage. +`kv_update` preserves each field's tensor/non-tensor type and tensor dtype. A +timeout has an unknown outcome because a storage unit may commit after the +caller stops waiting, so non-idempotent merge operations must not be retried +blindly. Updates are atomic within one storage-unit request; a batch spanning +multiple units is not a distributed transaction. + **Key Features** - **Redis-style Semantics**: Familiar KV interface (Put/Get/List) for a zero learning curve. diff --git a/docs/ssd_offload.md b/docs/ssd_offload.md index b0e5b66a..99507bfb 100644 --- a/docs/ssd_offload.md +++ b/docs/ssd_offload.md @@ -213,10 +213,22 @@ GET reads each requested SSD file and reconstructs the value in host memory. SSD offload reduces long-lived memory use, but it does not remove temporary memory use while values are read or encoded. -### `data_parser` outputs should not share backing storage across samples +### Dense batches can retain shared backing memory -If a `data_parser` returns differently sized tensor or NumPy views backed by -the same allocation, an in-memory view can keep the entire allocation resident -after another view is offloaded. When parsing URLs or file paths, return an -independently owned value for each sample, or copy retained views before -returning them. +SimpleStorage slices dense batches without copying. The rows therefore share +the batch allocation, and replacing or emptying one row does not release its +share while another stored row still references that allocation. This is the +memory trade-off for zero-copy batch distribution; clear or replace every +sibling row before expecting the full allocation to be released. + +The same rule applies when a `data_parser` returns tensor or NumPy views backed +by one allocation. Return independently owned values when partial release is +more important than zero-copy storage. + +### Atomic multi-field writes require temporary SSD headroom + +SimpleStorage writes all replacement files before publishing any of them. This +keeps a failed multi-field put or merge from exposing a partially updated +sample, but temporarily requires space for the old files and all new files. +When sizing the SSD tier, allow for that operation-level peak rather than only +the final active-byte count. diff --git a/tests/e2e/test_kv_interface_e2e.py b/tests/e2e/test_kv_interface_e2e.py index 271758ee..875aedaf 100644 --- a/tests/e2e/test_kv_interface_e2e.py +++ b/tests/e2e/test_kv_interface_e2e.py @@ -1110,6 +1110,156 @@ def test_field_expansion_across_samples(self, controller, tq_api): tq_api.kv_clear(keys=keys, partition_id=partition_id) +class TestKVUpdateE2E: + """SimpleStorage merge and empty behavior through sync and async public APIs.""" + + def test_update_then_empty_preserves_partition_field_metadata(self, controller, tq_api, backend_name): + if backend_name != "SimpleStorage": + pytest.skip("merge-backed kv_update is implemented only for SimpleStorage") + + partition_id = "test_partition" + keys = ["empty_0", "empty_1", "empty_2"] + tq_api.kv_batch_put( + keys=keys, + partition_id=partition_id, + fields=TensorDict({"tokens": torch.arange(12).reshape(3, 4)}, batch_size=3), + ) + + def concat(old, new): + return torch.cat([old, new]) + + meta = tq_api.kv_update( + key=keys[1], + partition_id=partition_id, + fields={"tokens": torch.tensor([12, 13])}, + merge_fn=concat, + ) + assert "tokens" in meta.fields + tq_api.kv_empty(keys=keys[0], partition_id=partition_id, fields="tokens") + emptied = tq_api.kv_batch_get(keys=keys[0], partition_id=partition_id, select_fields="tokens") + assert emptied["tokens"][0] is None + + partition = get_controller_partition(controller, partition_id) + col = partition.field_name_mapping["tokens"] + global_idx = partition.keys_mapping[keys[0]] + assert partition.production_status[global_idx, col] == 1 + assert partition.field_metadata["tokens"].is_non_tensor is False + assert partition.field_metadata["tokens"].dtype == torch.int64 + + tq_api.kv_batch_put( + keys=["nested_0", "nested_1"], + partition_id=partition_id, + fields=TensorDict( + { + "tokens": torch.nested.as_nested_tensor( + [torch.tensor([20, 21]), torch.tensor([30, 31, 32])], + layout=torch.jagged, + ) + }, + batch_size=2, + ), + ) + partition = get_controller_partition(controller, partition_id) + assert partition.field_metadata["tokens"].is_nested is True + + def test_batch_update_multiple_fields_across_storage_units(self, controller, tq_api, backend_name): + if backend_name != "SimpleStorage": + pytest.skip("merge-backed kv_update is implemented only for SimpleStorage") + + partition_id = "test_partition" + keys = [f"batch_update_{i}" for i in range(8)] + lengths = [i + 1 for i in range(len(keys))] + tq_api.kv_batch_put( + keys=keys, + partition_id=partition_id, + fields=TensorDict( + { + "tokens": torch.nested.as_nested_tensor( + [torch.arange(length) for length in lengths], + layout=torch.jagged, + ), + "score": torch.arange(8).reshape(8, 1), + }, + batch_size=8, + ), + ) + tq_api.kv_batch_update( + keys=keys, + partition_id=partition_id, + fields=TensorDict( + { + "tokens": torch.arange(100, 108).reshape(8, 1), + "score": torch.arange(200, 208).reshape(8, 1), + }, + batch_size=8, + ), + merge_fn=lambda old, new: torch.cat([old, new]), + ) + + retrieved = tq_api.kv_batch_get(keys=keys, partition_id=partition_id) + for i, length in enumerate(lengths): + assert_tensor_equal( + retrieved["tokens"][i], + torch.cat([torch.arange(length), torch.tensor([100 + i])]), + ) + assert_tensor_equal(retrieved["score"][i], torch.tensor([i, 200 + i])) + + partition = get_controller_partition(controller, partition_id) + tokens = partition.field_metadata["tokens"] + assert tokens.is_nested is True + for key, length in zip(keys, lengths, strict=True): + global_idx = partition.keys_mapping[key] + assert tuple(tokens.per_sample_shapes[global_idx]) == (length + 1,) + + tq_api.kv_empty(keys=keys, partition_id=partition_id, fields=["tokens", "score"]) + emptied = tq_api.kv_batch_get(keys=keys, partition_id=partition_id) + assert all(value is None for value in emptied["tokens"]) + assert all(value is None for value in emptied["score"]) + + def test_rejects_missing_key_unproduced_field_and_incompatible_result(self, tq_api, backend_name): + if backend_name != "SimpleStorage": + pytest.skip("merge-backed kv_update is implemented only for SimpleStorage") + + with pytest.raises(ValueError, match="not found"): + tq_api.kv_update( + key="no_such", + partition_id="test_partition", + fields={"tokens": torch.tensor([1])}, + merge_fn=lambda _old, new: new, + ) + tq_api.kv_put(key="sample", partition_id="test_partition", fields={"tokens": torch.tensor([1, 2])}) + with pytest.raises(ValueError, match="not produced"): + tq_api.kv_update( + key="sample", + partition_id="test_partition", + fields={"fresh": torch.tensor([1])}, + merge_fn=lambda _old, new: new, + ) + with pytest.raises(RuntimeError, match="keep dtype"): + tq_api.kv_update( + key="sample", + partition_id="test_partition", + fields={"tokens": torch.tensor([3])}, + merge_fn=lambda old, new: torch.cat([old, new]).to(torch.float32), + ) + unchanged = tq_api.kv_batch_get(keys="sample", partition_id="test_partition", select_fields="tokens") + assert_tensor_equal(unchanged["tokens"][0], torch.tensor([1, 2])) + + def test_kv_backend_rejects(self, tq_api, backend_name): + if backend_name == "SimpleStorage": + pytest.skip("rejection is the KV-backend contract") + + tq_api.kv_put(key="k", partition_id="test_partition", fields={"tokens": torch.tensor([1])}) + with pytest.raises(NotImplementedError, match="kv_update is not supported"): + tq_api.kv_update( + key="k", + partition_id="test_partition", + fields={"tokens": torch.tensor([2])}, + merge_fn=lambda _old, new: new, + ) + tq_api.kv_clear(keys="k", partition_id="test_partition") + + def run_tests(): """Run all e2e tests manually for debugging.""" pytest.main([__file__, "-v", "-s"]) diff --git a/tests/test_kv_update.py b/tests/test_kv_update.py new file mode 100644 index 00000000..567722ff --- /dev/null +++ b/tests/test_kv_update.py @@ -0,0 +1,219 @@ +# Copyright 2025 Huawei Technologies Co., Ltd. All Rights Reserved. +# Copyright 2025 The TransferQueue Team +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import errno +from unittest.mock import MagicMock, patch + +import numpy as np +import pytest +import torch +from tensordict import TensorDict + +from transfer_queue.interface import _single_update_batch +from transfer_queue.metadata import BatchMeta +from transfer_queue.storage.managers.base import KVStorageManager +from transfer_queue.storage.managers.mooncake_manager import MooncakeStorageManager +from transfer_queue.storage.managers.ray_storage_manager import RayStorageManager +from transfer_queue.storage.managers.simple_storage_manager import _build_update_field_schema +from transfer_queue.storage.managers.yuanrong_manager import YuanrongStorageManager +from transfer_queue.storage.simple_storage import HybridStorageUnitData, StorageUnitData + + +def _concat(old, new): + return torch.cat([old, new]) + + +def _tensor_schema(*fields, dtype=torch.int64): + return {field: {"dtype": dtype, "shape": None, "is_nested": True, "is_non_tensor": False} for field in fields} + + +def test_single_update_batch_wraps_tensor_and_python_values(): + batch = _single_update_batch({"tokens": torch.tensor([4, 5]), "meta": {"step": 1}}) + + assert batch.batch_size == torch.Size([1]) + assert torch.equal(batch["tokens"][0], torch.tensor([4, 5])) + assert batch["meta"][0] == {"step": 1} + + +@pytest.mark.parametrize("fields", [{}, [], {"tokens": torch.tensor([1]), 2: "bad"}]) +def test_single_update_batch_rejects_invalid_fields(fields): + with pytest.raises(TypeError): + _single_update_batch(fields) + + +def test_apply_update_concatenates_and_keeps_field_name(): + data = StorageUnitData() + data.put_data({"sequence_ids": [torch.tensor([10, 11, 12])]}, [0]) + + described = data.apply_update( + [0], + {"sequence_ids": [torch.tensor([20, 21])]}, + _concat, + _tensor_schema("sequence_ids"), + ) + + assert described["sequence_ids"] == {"dtype": torch.int64, "shapes": [(5,)]} + assert torch.equal(data.field_data["sequence_ids"][0], torch.tensor([10, 11, 12, 20, 21])) + assert set(data.field_data) == {"sequence_ids"} + + +def test_apply_update_is_atomic_on_merge_error(): + data = StorageUnitData() + original = torch.tensor([1, 2, 3]) + data.put_data({"tokens": [original.clone()]}, [7]) + + def fail(_old, _new): + raise RuntimeError("merge failed") + + with pytest.raises(RuntimeError, match="merge failed"): + data.apply_update([7], {"tokens": [torch.tensor([9])]}, fail, _tensor_schema("tokens")) + assert torch.equal(data.field_data["tokens"][7], original) + + +def test_apply_update_rejects_unproduced_field(): + data = StorageUnitData() + data.put_data({"tokens": [torch.tensor([1])]}, [3]) + + with pytest.raises(ValueError, match="unproduced field"): + data.apply_update([3], {"fresh": [torch.tensor([2])]}, lambda _old, new: new, _tensor_schema("fresh")) + assert "fresh" not in data.field_data + + +@pytest.mark.parametrize( + "merge_fn, error", + [ + (lambda old, _new: old.to(torch.float32), "keep dtype"), + (lambda _old, _new: {"not": "a tensor"}, "keep dtype"), + ], +) +def test_apply_update_rejects_incompatible_tensor_results_before_writing(merge_fn, error): + data = StorageUnitData() + original = torch.tensor([1, 2]) + data.put_data({"tokens": [original.clone()]}, [0]) + + with pytest.raises(TypeError, match=error): + data.apply_update([0], {"tokens": [torch.tensor([3])]}, merge_fn, _tensor_schema("tokens")) + assert torch.equal(data.field_data["tokens"][0], original) + + +def test_apply_update_rejects_tensor_result_for_non_tensor_field(): + data = StorageUnitData() + data.put_data({"meta": [{"step": 1}]}, [0]) + schema = {"meta": {"dtype": None, "shape": None, "is_nested": False, "is_non_tensor": True}} + + with pytest.raises(TypeError, match="must remain non-tensor"): + data.apply_update([0], {"meta": [{"step": 2}]}, lambda _old, _new: torch.tensor([2]), schema) + assert data.field_data["meta"][0] == {"step": 1} + + +def test_apply_update_decodes_ssd_offloaded_old_value(tmp_path): + data = HybridStorageUnitData( + storage_size=4, threshold_bytes=64, ssd_path=str(tmp_path), run_id="run", unit_id="unit" + ) + prompt = torch.arange(32) + data.put_data({"tokens": [prompt]}, [0]) + assert data.ssd_active_values == 1 + + data.apply_update([0], {"tokens": [torch.tensor([99])]}, _concat, _tensor_schema("tokens")) + + assert torch.equal(data.get_data(["tokens"], [0])["tokens"][0], torch.cat([prompt, torch.tensor([99])])) + assert data.ssd_active_values == 1 + + +def test_apply_update_is_atomic_when_a_later_field_fails_to_reach_ssd(tmp_path, monkeypatch): + data = HybridStorageUnitData( + storage_size=4, threshold_bytes=64, ssd_path=str(tmp_path), run_id="run", unit_id="unit" + ) + try: + prompt, mask = torch.arange(32), torch.ones(32, dtype=torch.int64) + data.put_data({"tokens": [prompt], "mask": [mask]}, [0]) + old_files = set(tmp_path.rglob("*.bin")) + write_values = data._ssd_store.write_values + calls = [] + + def fail_second_field(encoded_values): + calls.append(encoded_values) + if len(calls) == 2: + raise OSError(errno.ENOSPC, "No space left on device") + return write_values(encoded_values) + + monkeypatch.setattr(data._ssd_store, "write_values", fail_second_field) + new_data = {"tokens": [torch.tensor([99])], "mask": [torch.tensor([1])]} + with pytest.raises(OSError, match="No space left"): + data.apply_update([0], new_data, _concat, _tensor_schema("tokens", "mask")) + + stored = data.get_data(["tokens", "mask"], [0]) + assert torch.equal(stored["tokens"][0], prompt) + assert torch.equal(stored["mask"][0], mask) + assert set(tmp_path.rglob("*.bin")) == old_files + assert (data.ssd_active_values, data.ssd_active_bytes) == (2, prompt.nbytes + mask.nbytes) + finally: + data.close() + + +def test_build_update_field_schema_orders_shapes_across_units(): + described = _build_update_field_schema( + [0, 1, 2, 3], + [ + ([0, 2], {"tokens": {"dtype": torch.int64, "shapes": [(4,), (4,)]}}), + ([1, 3], {"tokens": {"dtype": torch.int64, "shapes": [(9,), (9,)]}}), + ], + ) + + assert described["tokens"]["is_nested"] is True + assert described["tokens"]["shape"] is None + assert described["tokens"]["per_sample_shapes"] == [(4,), (9,), (4,), (9,)] + + +def test_build_update_field_schema_keeps_uniform_column_flat(): + described = _build_update_field_schema( + [0, 1], + [ + ([0], {"tokens": {"dtype": torch.int64, "shapes": [(4,)]}}), + ([1], {"tokens": {"dtype": torch.int64, "shapes": [(4,)]}}), + ], + ) + + assert described["tokens"] == { + "dtype": torch.int64, + "shape": (4,), + "is_nested": False, + "is_non_tensor": False, + } + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "manager_cls", [MooncakeStorageManager, YuanrongStorageManager, RayStorageManager, KVStorageManager] +) +@patch("transfer_queue.storage.managers.base.StorageClientFactory.create") +@patch.object(KVStorageManager, "_connect_to_controller", lambda self: None) +async def test_kv_backends_reject_update(mock_create, manager_cls): + mock_create.return_value = MagicMock() + config = { + KVStorageManager: {"client_name": "YuanrongStorageClient"}, + YuanrongStorageManager: {"worker_port": 31501}, + }.get(manager_cls, {}) + manager = manager_cls(controller_info=MagicMock(), config=config) + meta = BatchMeta( + global_indexes=[0], + partition_ids=["p"], + field_schema={"x": {"dtype": torch.int64, "shape": (1,), "is_nested": False, "is_non_tensor": False}}, + production_status=np.ones(1, dtype=np.int8), + ) + values = TensorDict({"x": torch.ones(1, 1, dtype=torch.int64)}, batch_size=1) + + with pytest.raises(NotImplementedError, match=f"not supported by {manager_cls.__name__}"): + await manager.update_data(meta, values, lambda _old, new: new) diff --git a/tests/test_simple_storage_unit.py b/tests/test_simple_storage_unit.py index 99ceb4ec..fe52f175 100644 --- a/tests/test_simple_storage_unit.py +++ b/tests/test_simple_storage_unit.py @@ -28,6 +28,7 @@ # See the License for the specific language governing permissions and # limitations under the License. +import errno import pickle import time from pathlib import Path @@ -603,6 +604,39 @@ def fail_write(_fd, _data): storage.close() +def test_failed_ssd_write_on_second_field_leaves_every_field_unchanged(tmp_path, monkeypatch): + storage = HybridStorageUnitData( + storage_size=10, + threshold_bytes=64, + ssd_path=str(tmp_path), + run_id="test-run", + unit_id="test-unit", + ) + try: + storage.put_data({"a": [b"a" * 100], "b": [b"b" * 100]}, [1]) + old_files = set(tmp_path.rglob("*.bin")) + + write_values = storage._ssd_store.write_values + calls = [] + + def fail_second_field(encoded_values): + calls.append(encoded_values) + if len(calls) == 2: + raise OSError(errno.ENOSPC, "No space left on device") + return write_values(encoded_values) + + monkeypatch.setattr(storage._ssd_store, "write_values", fail_second_field) + with pytest.raises(OSError, match="No space left"): + storage.put_data({"a": [b"A" * 100], "b": [b"B" * 100]}, [1]) + + assert storage.get_data(["a", "b"], [1]) == {"a": [b"a" * 100], "b": [b"b" * 100]} + assert set(tmp_path.rglob("*.bin")) == old_files + assert storage.ssd_active_values == 2 + assert storage.ssd_active_bytes == 200 + finally: + storage.close() + + @pytest.mark.parametrize("threshold_bytes", [0, -1]) def test_hybrid_storage_requires_positive_threshold(tmp_path, threshold_bytes): with pytest.raises(ValueError, match="threshold must be greater than zero"): diff --git a/tests/test_storage_request_retry.py b/tests/test_storage_request_retry.py index 59f82edc..70be8985 100644 --- a/tests/test_storage_request_retry.py +++ b/tests/test_storage_request_retry.py @@ -29,7 +29,7 @@ from transfer_queue.storage.managers.simple_storage_manager import AsyncSimpleStorageManager, StorageUnitTimeout from transfer_queue.utils import common from transfer_queue.utils.enum_utils import Role -from transfer_queue.utils.zmq_utils import ZMQServerInfo +from transfer_queue.utils.zmq_utils import ZMQMessage, ZMQRequestType, ZMQServerInfo def _manager(with_unit: bool = True) -> AsyncSimpleStorageManager: @@ -57,6 +57,37 @@ def _single_sample_batch() -> tuple[TensorDict, BatchMeta]: return TensorDict({"input_ids": torch.zeros(1, 2, dtype=torch.int64)}, batch_size=1), metadata +class _SocketLease: + def __init__(self, socket): + self.socket = socket + + async def __aenter__(self): + return self.socket + + async def __aexit__(self, *_args): + return None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("accepted", [False, True]) +async def test_controller_notification_returns_ack_status(accepted): + manager = _manager() + manager.controller_info = MagicMock(id="controller") + socket = MagicMock() + socket.send_multipart = AsyncMock() + socket.recv_multipart = AsyncMock( + return_value=ZMQMessage.create( + request_type=ZMQRequestType.NOTIFY_DATA_UPDATE_ACK, + sender_id="controller", + body={"success": accepted}, + ).serialize() + ) + manager.notify_pool = MagicMock() + manager.notify_pool.alease.return_value = _SocketLease(socket) + + assert await manager._notify_and_wait([b"request"]) is accepted + + @pytest.mark.asyncio @pytest.mark.parametrize("data_parser, expected_attempts", [(None, 3), (lambda data: data, 1)]) async def test_put_retries_only_without_a_parser(data_parser, expected_attempts): @@ -76,6 +107,36 @@ async def test_put_retries_only_without_a_parser(data_parser, expected_attempts) manager.notify_data_update.assert_not_awaited() +@pytest.mark.asyncio +async def test_update_reports_controller_rejection_after_storage_commit(): + manager = _manager() + manager._update_to_single_storage_unit = AsyncMock( + return_value={"input_ids": {"dtype": torch.int64, "shapes": [(3,)]}} + ) + manager.notify_data_update = AsyncMock(return_value=False) + values, metadata = _single_sample_batch() + + with pytest.raises(RuntimeError, match="Storage update committed.*do not retry"): + await manager.update_data(metadata, values, lambda old, new: torch.cat([old, new])) + + manager._update_to_single_storage_unit.assert_awaited_once() + manager.notify_data_update.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_update_timeout_is_not_retried_or_published(): + manager = _manager() + manager._update_to_single_storage_unit = AsyncMock(side_effect=StorageUnitTimeout("outcome unknown")) + manager.notify_data_update = AsyncMock() + values, metadata = _single_sample_batch() + + with patch.object(manager, "_diagnose_storage_unit", return_value="diagnosis"), pytest.raises(StorageUnitTimeout): + await manager.update_data(metadata, values, lambda old, new: torch.cat([old, new])) + + manager._update_to_single_storage_unit.assert_awaited_once() + manager.notify_data_update.assert_not_awaited() + + @pytest.mark.asyncio async def test_lost_request_recovers_without_diagnosis(): manager = _manager() diff --git a/transfer_queue/__init__.py b/transfer_queue/__init__.py index 754bb4d8..350e9f53 100644 --- a/transfer_queue/__init__.py +++ b/transfer_queue/__init__.py @@ -21,9 +21,12 @@ async_kv_batch_get, async_kv_batch_get_by_meta, async_kv_batch_put, + async_kv_batch_update, async_kv_clear, + async_kv_empty, async_kv_list, async_kv_put, + async_kv_update, close, get_client, get_metrics_endpoint, @@ -31,9 +34,12 @@ kv_batch_get, kv_batch_get_by_meta, kv_batch_put, + kv_batch_update, kv_clear, + kv_empty, kv_list, kv_put, + kv_update, load_checkpoint, save_checkpoint, ) @@ -53,16 +59,22 @@ "get_metrics_endpoint", "kv_put", "kv_batch_put", + "kv_batch_update", "kv_batch_get", "kv_batch_get_by_meta", "kv_list", "kv_clear", + "kv_update", + "kv_empty", "async_kv_put", "async_kv_batch_put", + "async_kv_batch_update", "async_kv_batch_get", "async_kv_batch_get_by_meta", "async_kv_list", "async_kv_clear", + "async_kv_update", + "async_kv_empty", "KVBatchMeta", ] + [ diff --git a/transfer_queue/client.py b/transfer_queue/client.py index 0f7f13f1..d4c4d2bb 100644 --- a/transfer_queue/client.py +++ b/transfer_queue/client.py @@ -492,6 +492,36 @@ async def async_put( return metadata + async def async_update( + self, + metadata: BatchMeta, + values: TensorDict, + merge_fn: Callable[[Any, Any], Any], + ) -> BatchMeta: + """Merge new values into produced fields on SimpleStorage units. + + Args: + metadata: Samples to update. The key must already exist. + values: New values aligned with ``metadata``. + merge_fn: Called per sample per field as ``merge_fn(old, new)``. + + Returns: + The same metadata with field_schema and production status updated. + + Raises: + NotImplementedError: If the storage backend is not SimpleStorage. + """ + if not hasattr(self, "storage_manager") or self.storage_manager is None: + raise RuntimeError( + f"[{self.client_id}]: Storage manager not initialized. " + "Call initialize_storage_manager() before performing storage operations." + ) + if not metadata or metadata.size == 0: + raise ValueError("metadata cannot be none or empty") + + field_schema = await self.storage_manager.update_data(metadata, values, merge_fn) + return metadata.apply_field_schema(field_schema) + async def async_get_data(self, metadata: BatchMeta) -> TensorDict: """Asynchronously fetch data from storage units and organize into TensorDict. @@ -1294,6 +1324,7 @@ def wrapper(*args, **kwargs): # Bind internal sync wrappers. Public methods are defined explicitly below # to ensure proper type hints and documentation. self._put = _make_sync(self.async_put) + self._update = _make_sync(self.async_update) self._get_meta = _make_sync(self.async_get_meta) self._get_data = _make_sync(self.async_get_data) self._clear_partition = _make_sync(self.async_clear_partition) @@ -1480,6 +1511,18 @@ def put( """ return self._put(data=data, metadata=metadata, partition_id=partition_id, data_parser=data_parser) + def update( + self, + metadata: BatchMeta, + values: TensorDict, + merge_fn: Callable[[Any, Any], Any], + ) -> BatchMeta: + """Synchronously merge new values into produced fields on SimpleStorage units. + + See ``async_update``. + """ + return self._update(metadata, values, merge_fn) + def get_data(self, metadata: BatchMeta) -> TensorDict: """Synchronously fetch data from storage units and organize into TensorDict. diff --git a/transfer_queue/interface.py b/transfer_queue/interface.py index eca8762e..d70c3a60 100644 --- a/transfer_queue/interface.py +++ b/transfer_queue/interface.py @@ -606,6 +606,67 @@ def kv_clear(keys: list[str] | str, partition_id: str) -> None: tq_client._run_coroutine(async_kv_clear(keys=keys, partition_id=partition_id)) +def kv_update( + key: str, + partition_id: str, + fields: dict[str, Any], + merge_fn: Callable[[Any, Any], Any], +) -> KVBatchMeta: + """Merge new values into produced fields of one existing key. + + Args: + key: Existing user-specified key. The key is not created if it is missing. + partition_id: Partition that holds the key. + fields: New values keyed by an already-produced field name. + merge_fn: ``merge_fn(old, new) -> stored``. SimpleStorage runs it once + per field on the unit that owns the sample. + + Returns: + Metadata for the updated sample. + + Raises: + ValueError: If the key or a field is missing. + TypeError: If ``fields`` is not a non-empty dict or ``merge_fn`` is not callable. + NotImplementedError: If the storage backend is not SimpleStorage. + + Note: + A timeout has an unknown outcome because the storage unit may finish the + merge after the caller stops waiting. Do not blindly retry a merge. + """ + tq_client = _maybe_create_tq_client() + return tq_client._run_coroutine( + async_kv_update( + key=key, + partition_id=partition_id, + fields=fields, + merge_fn=merge_fn, + ) + ) + + +def kv_batch_update( + keys: list[str], + partition_id: str, + fields: TensorDict, + merge_fn: Callable[[Any, Any], Any], +) -> KVBatchMeta: + """Merge batched values into produced fields of existing keys. + + A timeout has an unknown outcome because a unit may commit after the caller + stops waiting. Do not blindly retry a merge. + """ + tq_client = _maybe_create_tq_client() + return tq_client._run_coroutine( + async_kv_batch_update(keys=keys, partition_id=partition_id, fields=fields, merge_fn=merge_fn) + ) + + +def kv_empty(keys: str | list[str], partition_id: str, fields: str | list[str]) -> KVBatchMeta: + """Release produced fields by storing ``None`` while keeping the keys ready.""" + tq_client = _maybe_create_tq_client() + return tq_client._run_coroutine(async_kv_empty(keys=keys, partition_id=partition_id, fields=fields)) + + # ==================== KV Interface API ==================== async def async_kv_put( key: str, @@ -987,6 +1048,110 @@ async def async_kv_clear(keys: list[str] | str, partition_id: str) -> None: await tq_client.async_clear_samples(batch_meta) +def _single_update_batch(fields: dict[str, Any]) -> TensorDict: + """Wrap one sample's field mapping in a one-row TensorDict.""" + if not isinstance(fields, dict) or not fields: + raise TypeError("fields must be a non-empty dict") + batch: dict[str, Any] = {} + for field_name, value in fields.items(): + if not isinstance(field_name, str): + raise TypeError("field names must be strings") + if isinstance(value, torch.Tensor): + if value.is_nested: + raise ValueError("Use async_kv_batch_update for nested tensors") + batch[field_name] = value.unsqueeze(0) + else: + batch[field_name] = NonTensorStack(value) + return TensorDict(batch, batch_size=[1]) + + +async def async_kv_update( + key: str, + partition_id: str, + fields: dict[str, Any], + merge_fn: Callable[[Any, Any], Any], +) -> KVBatchMeta: + """Asynchronously update fields of an existing key. See ``kv_update``.""" + return await async_kv_batch_update( + keys=[key], + partition_id=partition_id, + fields=_single_update_batch(fields), + merge_fn=merge_fn, + ) + + +async def async_kv_batch_update( + keys: list[str], + partition_id: str, + fields: TensorDict, + merge_fn: Callable[[Any, Any], Any], +) -> KVBatchMeta: + """Asynchronously merge batched values into produced fields of existing keys.""" + if not isinstance(keys, list) or not keys or not all(isinstance(key, str) for key in keys): + raise TypeError("keys must be a non-empty list of strings") + if not isinstance(fields, TensorDict) or not fields.keys(): + raise TypeError("fields must be a non-empty TensorDict") + if fields.batch_size != torch.Size([len(keys)]): + raise ValueError(f"fields batch size {fields.batch_size} does not match {len(keys)} keys") + if not callable(merge_fn): + raise TypeError("merge_fn must be callable") + tq_client = _maybe_create_tq_client() + batch_meta = await tq_client.async_kv_retrieve_meta(keys=keys, partition_id=partition_id, create=False) + + if batch_meta.size != len(keys): + raise ValueError("Some keys or the partition were not found") + missing_fields = sorted(set(fields.keys()) - set(batch_meta.field_names)) + if missing_fields: + raise ValueError(f"Fields are not produced for every key: {missing_fields}") + + batch_meta = await tq_client.async_update(metadata=batch_meta, values=fields, merge_fn=merge_fn) + + return KVBatchMeta( + keys=keys, + tags=batch_meta.custom_meta, + partition_id=partition_id, + fields=batch_meta.field_names, + extra_info=batch_meta.extra_info, + ) + + +async def async_kv_empty(keys: str | list[str], partition_id: str, fields: str | list[str]) -> KVBatchMeta: + """Asynchronously release fields by storing ``None`` through the ordinary put path.""" + if isinstance(keys, str): + keys = [keys] + if isinstance(fields, str): + field_names = [fields] + elif isinstance(fields, list) and fields and all(isinstance(name, str) for name in fields): + field_names = list(dict.fromkeys(fields)) + else: + raise TypeError("fields must be a field name or a non-empty list of field names") + if not keys: + raise ValueError("keys must not be empty") + + tq_client = _maybe_create_tq_client() + if not isinstance(tq_client.storage_manager, AsyncSimpleStorageManager): + raise NotImplementedError("kv_empty requires the SimpleStorage backend") + batch_meta = await tq_client.async_kv_retrieve_meta(keys=keys, partition_id=partition_id, create=False) + if batch_meta.size != len(keys): + raise ValueError("Some keys or the partition were not found") + missing_fields = sorted(set(field_names) - set(batch_meta.field_names)) + if missing_fields: + raise ValueError(f"Fields are not produced for every key: {missing_fields}") + + empty_values = TensorDict( + {field_name: NonTensorStack(*([None] * len(keys))) for field_name in field_names}, + batch_size=[len(keys)], + ) + batch_meta = await tq_client.async_put(empty_values, batch_meta) + return KVBatchMeta( + keys=keys, + tags=batch_meta.custom_meta, + partition_id=partition_id, + fields=batch_meta.field_names, + extra_info=batch_meta.extra_info, + ) + + # ==================== Low-Level Native API ==================== # For low-level API support, please refer to transfer_queue/client.py for details. def get_client(): diff --git a/transfer_queue/metadata.py b/transfer_queue/metadata.py index e89c1531..ab395e07 100644 --- a/transfer_queue/metadata.py +++ b/transfer_queue/metadata.py @@ -461,10 +461,20 @@ def add_fields(self, tensor_dict: TensorDict, set_all_ready: bool = True) -> "Ba if batch_size != self.size: raise ValueError(f"add_fields batch size mismatch: self.size={self.size} vs tensor_dict={batch_size}") - field_schema = extract_field_schema(tensor_dict) + return self.apply_field_schema(extract_field_schema(tensor_dict), set_all_ready=set_all_ready) - for key, value in field_schema.items(): - self.field_schema[key] = value + def apply_field_schema(self, field_schema: dict[str, dict[str, Any]], set_all_ready: bool = True) -> "BatchMeta": + """Merge a field_schema into this batch in place. + + The only place that keeps ``field_names`` and ``is_ready`` in step with + ``field_schema`` and ``production_status``; callers holding a schema from storage + should go through here rather than writing the derived state themselves. + + Args: + field_schema: Per-field metadata to add or overwrite. + set_all_ready (bool): If True, set all production_status to READY_FOR_CONSUME. Default is True. + """ + self.field_schema.update(field_schema) if set_all_ready: self.production_status[:] = 1 diff --git a/transfer_queue/storage/managers/base.py b/transfer_queue/storage/managers/base.py index 369f800e..7c0fe6fe 100644 --- a/transfer_queue/storage/managers/base.py +++ b/transfer_queue/storage/managers/base.py @@ -226,7 +226,7 @@ async def notify_data_update( global_indexes: list[int], field_schema: dict[str, dict[str, Any]], custom_backend_meta: dict[int, dict[str, Any]] | None = None, - ) -> None: + ) -> bool: """ Notify controller that new data is ready. @@ -239,7 +239,7 @@ async def notify_data_update( if not self.controller_info: logger.warning(f"No controller connected for storage manager {self.storage_manager_id}") - return + return False normalized_field_schema = {} for field_name, field in field_schema.items(): @@ -271,9 +271,9 @@ async def notify_data_update( self._notify_and_wait(request_msg), self._notify_loop, ) - await asyncio.wrap_future(thread_future) + return await asyncio.wrap_future(thread_future) - async def _notify_and_wait(self, request_msg: list) -> None: + async def _notify_and_wait(self, request_msg: list) -> bool: """Send a data status notification to the controller and block until ACK is received.""" # Acquiring the lease sits outside the handler below: a missing socket name or a dead # context is a configuration/lifecycle fault the caller must see, not a slow ACK. @@ -298,16 +298,18 @@ async def _notify_and_wait(self, request_msg: list) -> None: response_msg = ZMQMessage.deserialize(messages) if response_msg.request_type == ZMQRequestType.NOTIFY_DATA_UPDATE_ACK: # type: ignore[arg-type] + success = bool(response_msg.body.get("success")) logger.debug( - f"[{self.storage_manager_id}]: Get data status update ACK response " - f"from controller id #{response_msg.sender_id} successfully." + f"[{self.storage_manager_id}]: Received data status update ACK " + f"from controller id #{response_msg.sender_id}: success={success}." ) - return + return success except Exception as e: # Logged rather than raised, so a slow controller does not fail the put. Close # the socket: a late ACK would otherwise be read as the next lessee's reply. logger.error(f"[{self.storage_manager_id}]: Data status update failed: {type(e).__name__}: {e}") sock.close(linger=0) + return False @abstractmethod async def put_data( @@ -354,6 +356,26 @@ async def clear_data(self, metadata: BatchMeta) -> None: """ raise NotImplementedError("Subclasses must implement clear_data") + async def update_data( + self, + metadata: BatchMeta, + values: TensorDict, + merge_fn: Callable[[Any, Any], Any], + ) -> dict[str, dict[str, Any]]: + """Merge new values on the unit that holds each sample. + + Args: + metadata: Samples to update. + values: New values as a TensorDict aligned with ``metadata``. + merge_fn: Called per sample per field as ``merge_fn(old, new)``. + + Returns: + field_schema of the stored values, keyed by field name. + """ + raise NotImplementedError( + f"kv_update is not supported by {self.__class__.__name__}; it requires the SimpleStorage backend" + ) + async def save_checkpoint(self, checkpoint_dir: str) -> None: """Save storage state into checkpoint_dir. diff --git a/transfer_queue/storage/managers/simple_storage_manager.py b/transfer_queue/storage/managers/simple_storage_manager.py index c89d1cfb..ea23b6ad 100644 --- a/transfer_queue/storage/managers/simple_storage_manager.py +++ b/transfer_queue/storage/managers/simple_storage_manager.py @@ -87,6 +87,48 @@ def _describe_unit_state(body: dict[str, Any]) -> str: return f"verdict=unit_serving_again ({' '.join(parts)})" +def _build_update_field_schema( + global_indexes: list[int], + per_unit: list[tuple[list[int], dict[str, dict[str, Any]]]], +) -> dict[str, dict[str, Any]]: + """Assemble one batch-level field_schema from each unit's per-sample description. + + A unit only describes its own rows, so shapes are collected per global index and + then ordered to match ``global_indexes``: the controller reads per_sample_shapes + positionally against the indexes it is notified about. + """ + dtypes: dict[str, Any] = {} + shapes: dict[str, dict[int, tuple]] = defaultdict(dict) + non_tensor: set[str] = set() + for unit_indexes, described in per_unit: + for field_name, description in described.items(): + unit_shapes = description["shapes"] + if unit_shapes is None: + non_tensor.add(field_name) + continue + dtypes[field_name] = description["dtype"] + # Serialization turns the unit's tuples into lists; shapes are compared and hashed here. + shapes[field_name].update(zip(unit_indexes, (tuple(s) for s in unit_shapes), strict=True)) + + field_schema: dict[str, dict[str, Any]] = {} + for field_name in non_tensor | set(shapes): + # A column is only a tensor column when every unit stored tensors for it. + if field_name in non_tensor: + field_schema[field_name] = {"dtype": None, "shape": None, "is_nested": False, "is_non_tensor": True} + continue + ordered = [shapes[field_name][gi] for gi in global_indexes] + is_nested = len(set(ordered)) > 1 + field_schema[field_name] = { + "dtype": dtypes[field_name], + "shape": None if is_nested else ordered[0], + "is_nested": is_nested, + "is_non_tensor": False, + } + if is_nested: + field_schema[field_name]["per_sample_shapes"] = ordered + return field_schema + + _SU_SUBDIR = "simple_storage" _SU_INFO_FILE = "storage_unit_info.json" @@ -481,6 +523,78 @@ async def put_data( field_schema, ) + async def update_data( + self, + metadata: BatchMeta, + values: TensorDict, + merge_fn: Callable[[Any, Any], Any], + ) -> dict[str, dict[str, Any]]: + """Merge values on each owner unit and notify the controller. + + Each unit reports the values it stored; this assembles one batch-level + field_schema so production status stays ready. + """ + logger.debug(f"[{self.storage_manager_id}]: receive update_data request, updating {metadata.size} samples.") + + if metadata.size == 0: + return {} + if values.batch_size[0] != metadata.size: + raise ValueError( + f"Batch size of values ({values.batch_size[0]}) does not match metadata size ({metadata.size})" + ) + field_names = list(values.keys()) + + routing = self._group_by_hash(metadata.global_indexes) + tasks = [] + for su_id, group in routing.items(): + storage_data = {f: self._select_by_positions(values[f], group.batch_positions) for f in field_names} + tasks.append( + self._request_with_retry( + "update", + su_id, + f"samples={len(group.global_indexes)} fields={field_names}", + partial( + self._update_to_single_storage_unit, + group.global_indexes, + storage_data, + merge_fn, + {name: metadata.field_schema[name] for name in field_names}, + target_storage_unit=su_id, + ), + # A merge reads stored state, so a replay after a lost answer + # could fold the same new value in twice. + max_attempts=1, + ) + ) + + try: + described = await asyncio.gather(*tasks) + except Exception as e: + logger.error( + f"[{self.storage_manager_id}]: update_data failed. " + f"partition_id={metadata.partition_ids[0]}, " + f"num_samples={metadata.size}, " + f"num_storage_units={len(routing)}, " + f"error={type(e).__name__}: {e}" + ) + raise + + field_schema = _build_update_field_schema( + metadata.global_indexes, list(zip([g.global_indexes for g in routing.values()], described, strict=True)) + ) + + published = await self.notify_data_update( + metadata.partition_ids[0], + metadata.global_indexes, + field_schema, + ) + if not published: + raise RuntimeError( + "Storage update committed, but the controller did not accept its metadata; " + "do not retry this merge blindly" + ) + return field_schema + @with_storage_unit_socket async def _put_to_single_storage_unit( self, @@ -538,6 +652,67 @@ async def _put_to_single_storage_unit( ) raise RuntimeError(f"Error in put to storage unit {target_storage_unit}: {type(e).__name__}: {e}") from e + @with_storage_unit_socket + async def _update_to_single_storage_unit( + self, + global_indexes: list[int], + storage_data: dict[str, Any], + merge_fn: Callable[[Any, Any], Any], + field_schema: dict[str, dict[str, Any]], + target_storage_unit: str, + socket: zmq.Socket = None, + ) -> dict[str, dict[str, Any]]: + """Send an update to one storage unit and return what it stored, per field.""" + body: dict[str, Any] = { + "global_indexes": global_indexes, + "data": storage_data, + "merge_fn": merge_fn, + "field_schema": field_schema, + } + + request_msg = ZMQMessage.create( + request_type=ZMQRequestType.UPDATE_DATA, # type: ignore[arg-type] + sender_id=self.storage_manager_id, + receiver_id=target_storage_unit, + body=body, + ) + + serialized_bytes = 0 + started = time.perf_counter() + try: + frames = request_msg.serialize() + serialized_bytes = sum(frame_nbytes(frame) or 0 for frame in frames) + await socket.send_multipart(frames, copy=False) + messages = await socket.recv_multipart(copy=False) + response_msg = ZMQMessage.deserialize(messages) + + if response_msg.request_type != ZMQRequestType.UPDATE_DATA_RESPONSE: + raise RuntimeError( + f"Failed to update data on storage unit {target_storage_unit}: " + f"{response_msg.body.get('message', 'Unknown error')}" + ) + log_heavy_operation( + self.storage_manager_id, + "update", + time.perf_counter() - started, + serialized_bytes, + f"to {target_storage_unit} at {self._describe_storage_unit(target_storage_unit)} " + f"samples={len(global_indexes)} fields={list(storage_data)}", + ) + return response_msg.body.get("stored_shapes", {}) + except zmq.error.Again as e: + raise StorageUnitTimeout( + f"no answer in {TQ_SIMPLE_STORAGE_SEND_RECV_TIMEOUT}s during update to storage unit " + f"{target_storage_unit} at {self._describe_storage_unit(target_storage_unit)}; " + f"serialized_mb={serialized_bytes / 2**20:.1f}" + ) from e + except Exception as e: + logger.error( + f"[{self.storage_manager_id}]: Unexpected error during update to storage unit " + f"{target_storage_unit}: {type(e).__name__}: {e}" + ) + raise RuntimeError(f"Error in update to storage unit {target_storage_unit}: {type(e).__name__}: {e}") from e + @staticmethod def _pack_field_values(values: list) -> torch.Tensor | NonTensorStack: """ diff --git a/transfer_queue/storage/simple_storage.py b/transfer_queue/storage/simple_storage.py index df28521e..2d0f59dc 100644 --- a/transfer_queue/storage/simple_storage.py +++ b/transfer_queue/storage/simple_storage.py @@ -38,7 +38,7 @@ from dataclasses import dataclass from pathlib import Path from threading import Event, Thread -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, Any, Callable from uuid import uuid4 import numpy as np @@ -110,6 +110,18 @@ class _SSDValueRef: TQ_ACCEPT_PROBE_INTERVAL = float(os.environ.get("TQ_ACCEPT_PROBE_INTERVAL", 0)) +def _describe_stored_values(values: list) -> dict[str, Any]: + """Describe per-sample stored values for the manager. + + A unit only holds part of a batch, so it reports one shape per sample and + leaves ``shapes=None`` for a column it cannot describe as tensors; only the + manager sees the whole batch and can turn this into a field_schema. + """ + if not all(isinstance(v, torch.Tensor) for v in values): + return {"dtype": None, "shapes": None} + return {"dtype": values[0].dtype, "shapes": [tuple(v.shape) for v in values]} + + class StorageKeyNotFoundError(KeyError): """Raised when a requested global index is absent from a storage unit. @@ -208,6 +220,7 @@ def put_data(self, field_data: dict[str, Any], global_indexes: list) -> None: f"StorageUnitData put_data: field '{f}' values length {len(values)} " f"!= global_indexes length {len(global_indexes)}, length mismatch" ) + for f, values in field_data.items(): if f not in self.field_data: self.field_data[f] = {} field_dict = self.field_data[f] @@ -215,6 +228,44 @@ def put_data(self, field_data: dict[str, Any], global_indexes: list) -> None: field_dict[key] = val self._active_keys.update(global_indexes) + def apply_update( + self, + global_indexes: list[int], + new_data: dict[str, Any], + merge_fn: Callable[[Any, Any], Any], + field_schema: dict[str, dict[str, Any]], + ) -> dict[str, dict[str, Any]]: + """Merge produced values, validate their types, then write once. + + All merge calls and type checks finish before one ``put_data`` call, so + a merge error or incompatible result leaves storage unchanged. + + Returns: + Per-sample description of the stored values, keyed by field name. + """ + computed: dict[str, list] = {} + for field, new_values in new_data.items(): + stored = self.field_data.get(field, {}) + missing = [idx for idx in global_indexes if idx not in stored] + if missing: + raise ValueError(f"Cannot update unproduced field {field!r} for indexes {missing[:20]}") + old_values = self.get_data([field], global_indexes)[field] + computed[field] = [merge_fn(old, new_values[i]) for i, old in enumerate(old_values)] + + expected = field_schema[field] + if expected.get("is_non_tensor"): + if any(isinstance(value, torch.Tensor) for value in computed[field]): + raise TypeError(f"Merge result for non-tensor field {field!r} must remain non-tensor") + else: + for value in computed[field]: + if not isinstance(value, torch.Tensor) or value.dtype != expected["dtype"]: + actual = value.dtype if isinstance(value, torch.Tensor) else type(value).__name__ + raise TypeError( + f"Merge result for tensor field {field!r} must keep dtype {expected['dtype']}, got {actual}" + ) + self.put_data(computed, global_indexes) + return {field: _describe_stored_values(values) for field, values in computed.items()} + def clear(self, keys: list[int]) -> None: """Remove data at given global index keys, immediately freeing memory. @@ -549,38 +600,47 @@ def put_data( unique_global_indexes = set(global_indexes) has_duplicate_indexes = len(unique_global_indexes) != len(global_indexes) - for field, values in field_data.items(): - prepared_values, entries, fallback_values = self._prepare_field_values(values) - + old_ssd_values = [] + for field in field_data: stored_field = self.field_data.get(field, {}) - old_ssd_values = [] for global_index in unique_global_indexes: old_value = stored_field.get(global_index) if isinstance(old_value, _SSDValueRef): old_ssd_values.append(old_value) - try: - super().put_data({field: prepared_values}, global_indexes) - except Exception: - for entry in entries: - self._ssd_store.unlink(entry) - raise - - obsolete_ssd_values = old_ssd_values - if has_duplicate_indexes: - retained_ssd_paths = set() + + # Write every field's files before replacing any value: kv_update merge functions are not + # idempotent, so a failure on a later field must not leave earlier fields committed. + prepared_data = {} + entries: list[_SSDValueRef] = [] + fallback_values = 0 + try: + for field, values in field_data.items(): + prepared_data[field], field_entries, field_fallback_values = self._prepare_field_values(values) + entries.extend(field_entries) + fallback_values += field_fallback_values + super().put_data(prepared_data, global_indexes) + except Exception: + for entry in entries: + self._ssd_store.unlink(entry) + raise + + obsolete_ssd_values = old_ssd_values + if has_duplicate_indexes: + retained_ssd_paths = set() + for field in field_data: for global_index in unique_global_indexes: retained_value = self.field_data[field][global_index] if isinstance(retained_value, _SSDValueRef): retained_ssd_paths.add(retained_value.path) - obsolete_ssd_values.extend(entry for entry in entries if entry.path not in retained_ssd_paths) + obsolete_ssd_values.extend(entry for entry in entries if entry.path not in retained_ssd_paths) - self._ssd_active_values += len(entries) - len(obsolete_ssd_values) - self._ssd_active_bytes += sum(entry.size_bytes for entry in entries) - sum( - value.size_bytes for value in obsolete_ssd_values - ) - self._ssd_fallback_values_total += fallback_values - for obsolete_value in obsolete_ssd_values: - self._ssd_store.unlink(obsolete_value) + self._ssd_active_values += len(entries) - len(obsolete_ssd_values) + self._ssd_active_bytes += sum(entry.size_bytes for entry in entries) - sum( + value.size_bytes for value in obsolete_ssd_values + ) + self._ssd_fallback_values_total += fallback_values + for obsolete_value in obsolete_ssd_values: + self._ssd_store.unlink(obsolete_value) def get_data(self, fields: list[str], global_indexes: list) -> dict[str, list]: """Read mixed memory- and SSD-backed samples in request order.""" @@ -993,6 +1053,9 @@ def _process_one_worker_request(self, worker_socket: zmq.Socket, monitor: Any) - if operation == ZMQRequestType.PUT_DATA: # type: ignore[arg-type] with monitor.measure(op_type="PUT_DATA"): response_msg = self._handle_put(request_msg) + elif operation == ZMQRequestType.UPDATE_DATA: # type: ignore[arg-type] + with monitor.measure(op_type="UPDATE_DATA"): + response_msg = self._handle_update(request_msg) elif operation == ZMQRequestType.GET_DATA: # type: ignore[arg-type] with monitor.measure(op_type="GET_DATA"): response_msg = self._handle_get(request_msg) @@ -1120,6 +1183,40 @@ def _handle_put(self, data_parts: ZMQMessage) -> ZMQMessage: }, ) + def _handle_update(self, data_parts: ZMQMessage) -> ZMQMessage: + """Merge existing fields, write once, and describe what was stored.""" + try: + global_indexes = data_parts.body["global_indexes"] + merge_fn = data_parts.body.get("merge_fn") + new_data = data_parts.body.get("data") + field_schema = data_parts.body.get("field_schema") + + with limit_pytorch_auto_parallel_threads( + target_num_threads=TQ_NUM_THREADS, info=f"[{self.storage_unit_id}] _handle_update" + ): + if not callable(merge_fn): + raise TypeError(f"merge_fn must be callable, got {type(merge_fn).__name__}") + if not isinstance(new_data, dict) or not new_data: + raise TypeError("update data must be a non-empty dict") + if not isinstance(field_schema, dict) or set(field_schema) != set(new_data): + raise ValueError("field_schema must describe every update field") + stored_shapes = self.storage_data.apply_update(global_indexes, new_data, merge_fn, field_schema) + + return ZMQMessage.create( + request_type=ZMQRequestType.UPDATE_DATA_RESPONSE, # type: ignore[arg-type] + sender_id=self.storage_unit_id, + body={"stored_shapes": stored_shapes}, + ) + except Exception as e: + return ZMQMessage.create( + request_type=ZMQRequestType.UPDATE_ERROR, # type: ignore[arg-type] + sender_id=self.storage_unit_id, + body={ + "message": f"Failed to update data in storage unit id " + f"#{self.storage_unit_id}, detail error message: {str(e)}" + }, + ) + def _handle_get(self, data_parts: ZMQMessage) -> ZMQMessage: """ Handle get request, return data from storage unit. @@ -1238,7 +1335,7 @@ def _handle_get_metrics(self) -> ZMQMessage: # Include per-operation stats if Prometheus metrics are enabled if self._metrics is not None: op_stats = {} - for op_type in ("PUT_DATA", "GET_DATA", "CLEAR_DATA"): + for op_type in ("PUT_DATA", "UPDATE_DATA", "GET_DATA", "CLEAR_DATA"): try: hist = self._metrics.request_duration.labels(op_type=op_type) counter = self._metrics.request_total.labels(op_type=op_type) diff --git a/transfer_queue/utils/zmq_utils.py b/transfer_queue/utils/zmq_utils.py index a94cd3ab..80742d8b 100644 --- a/transfer_queue/utils/zmq_utils.py +++ b/transfer_queue/utils/zmq_utils.py @@ -69,8 +69,10 @@ class ZMQRequestType(ExplicitEnum): # DATA_OPERATION GET_DATA = "GET" PUT_DATA = "PUT" + UPDATE_DATA = "UPDATE" GET_DATA_RESPONSE = "GET_DATA_RESPONSE" PUT_DATA_RESPONSE = "PUT_DATA_RESPONSE" + UPDATE_DATA_RESPONSE = "UPDATE_DATA_RESPONSE" CLEAR_DATA = "CLEAR_DATA" CLEAR_DATA_RESPONSE = "CLEAR_DATA_RESPONSE" @@ -78,6 +80,7 @@ class ZMQRequestType(ExplicitEnum): PUT_GET_ERROR = "PUT_GET_ERROR" PUT_ERROR = "PUT_ERROR" GET_ERROR = "GET_ERROR" + UPDATE_ERROR = "UPDATE_ERROR" CLEAR_DATA_ERROR = "CLEAR_DATA_ERROR" # META_OPERATION diff --git a/tutorial/02_kv_interface.py b/tutorial/02_kv_interface.py index 778739c2..bd182de1 100644 --- a/tutorial/02_kv_interface.py +++ b/tutorial/02_kv_interface.py @@ -56,10 +56,10 @@ def demonstrate_kv_api(): """ Demonstrate the Key-Value (KV) semantic API: - kv_put & kv_batch_put -> kv_list -> kv_batch_get -> kv_clear + kv_put & kv_batch_put -> kv_update -> kv_list -> kv_batch_get -> kv_clear """ print("=" * 80) - print("Key-Value Semantic API Demo: kv_put/kv_batch_put → kv_list → kv_batch_get → kv_clear") + print("Key-Value Semantic API Demo: kv_put/kv_batch_put → kv_update → kv_list → kv_batch_get → kv_clear") print("=" * 80) # Step 1: Put a single key-value pair with kv_put @@ -153,8 +153,21 @@ def demonstrate_kv_api(): tq.kv_put(key=key_for_update_tags, partition_id=partition_id, fields=None, tag=tag_update) print(f" ✓ Update success: Samples '0_0' now has tag as {tag_update}.") - # Step 5: List all keys and tags in a partition - print("\n[Step 5] Listing all keys and tags in partition...") + # Step 5: Concatenate a field with kv_update, then empty it + print("\n[Step 5] Updating a field with kv_update (concat, then empty)...") + tq.kv_put(key=key, partition_id=partition_id, fields={"scratch": torch.tensor([1, 2])}) + tq.kv_update( + key=key, + partition_id=partition_id, + fields={"scratch": torch.tensor([3, 4])}, + merge_fn=lambda old, new: torch.cat([old, new]), + ) + print(" ✓ kv_update concat: scratch of '0_0' is now [1, 2, 3, 4].") + tq.kv_empty(keys=key, partition_id=partition_id, fields="scratch") + print(" ✓ tq.kv_empty: scratch of '0_0' is stored as None (key remains).") + + # Step 6: List all keys and tags in a partition + print("\n[Step 6] Listing all keys and tags in partition...") partition_info = tq.kv_list() print(f" Found {len(partition_info.keys())} partitions: '{list(partition_info.keys())}'") @@ -162,16 +175,16 @@ def demonstrate_kv_api(): for k, t in keys_and_tags.items(): print(f"Partition: {pid}, - key='{k}' | tag={t}") - # Step 6: Retrieve specific fields using kv_batch_get - print("\n[Step 6] Retrieving specific fields (Column) with kv_batch_get...") + # Step 7: Retrieve specific fields using kv_batch_get + print("\n[Step 7] Retrieving specific fields (Column) with kv_batch_get...") print(" Fetching only 'input_ids' to save bandwidth (ignoring 'attention_mask' and 'response').") all_keys = list(partition_info[partition_id].keys()) retrieved_input_ids = tq.kv_batch_get(keys=all_keys, partition_id=partition_id, select_fields="input_ids") print(f" ✓ Successfully retrieved only {list(retrieved_input_ids.keys())} field for all samples.") - # # Step 7: Retrieve all fields using kv_batch_get - print("\n[Step 7] Retrieving all fields with kv_batch_get...") + # Step 8: Retrieve all fields using kv_batch_get + print("\n[Step 8] Retrieving all fields with kv_batch_get...") retrieved_all = tq.kv_batch_get(keys=all_keys, partition_id=partition_id) print(f" Retrieved all fields for {all_keys}:") print(f" Fields: {list(retrieved_all.keys())}") @@ -179,8 +192,8 @@ def demonstrate_kv_api(): f" Note: We cannot retrieve fields {list(response_batch.keys())}, since they only available in {append_keys}" ) - # Step 8: Clear specific keys - print("\n[Step 8] Clearing keys from partition...") + # Step 9: Clear specific keys + print("\n[Step 9] Clearing keys from partition...") keys_to_clear = all_keys[:2] # Delete the first 2 keys tq.kv_clear(keys=keys_to_clear, partition_id=partition_id) print(f" ✓ Cleared keys: {keys_to_clear}") @@ -202,9 +215,10 @@ def main(): Key Methods: 1. (async_)kv_put - Insert/Update a multi-column sample by key, with optional metadata tag 2. (async_)kv_batch_put - Put multiple key-value pairs efficiently in batch - 3. (async_)kv_batch_get - Retrieve samples (by keys), supporting column selection (by fields) - 4. (async_)kv_list - List keys and tags (metadata) in a partition - 5. (async_)kv_clear - Remove key-value pairs from storage + 3. (async_)kv_update - Merge new values into produced fields with merge_fn(old, new) + 4. (async_)kv_batch_get - Retrieve samples (by keys), supporting column selection (by fields) + 5. (async_)kv_list - List keys and tags (metadata) in a partition + 6. (async_)kv_clear - Remove key-value pairs from storage Key Features: ✓ Redis-style Semantics - Familiar KV interface (Put/Get/List) for zero learning curve