"""End-to-end tests for save_checkpoint and load_checkpoint.
Run with:
pytest tests/e2e/test_checkpoint_e2e.py -v
"""
import json
import os
import pytest
import ray
import torch
from omegaconf import OmegaConf
from tensordict import NonTensorStack, TensorDict
import transfer_queue as tq
os.environ["RAY_DEDUP_LOGS"] = "0"
_TQ_CONFIG = OmegaConf.create(
{
"controller": {"polling_mode": True},
"backend": {
"storage_backend": "SimpleStorage",
"SimpleStorage": {
"total_storage_size": 200,
"num_data_storage_units": 2,
},
},
}
)
@pytest.fixture(scope="module")
def ray_init():
if not ray.is_initialized():
ray.init(namespace="TestCheckpointE2E")
yield
if ray.is_initialized():
ray.shutdown()
@pytest.fixture(scope="module")
def tq_system(ray_init):
tq.init(_TQ_CONFIG)
yield
tq.close()
@pytest.fixture
def controller(tq_system):
return ray.get_actor("TransferQueueController", namespace="transfer_queue")
@pytest.fixture(autouse=True)
def cleanup_partitions(controller):
yield
try:
for pid in ray.get(controller.list_partitions.remote()):
ray.get(controller.clear_partition.remote(pid))
except Exception:
pass
@pytest.fixture
def checkpoint_dir(tmp_path):
return tmp_path / "checkpoint"
def _assert_tensor_equal(a, b, msg=""):
if (isinstance(a, torch.Tensor) and a.is_nested) or (isinstance(b, torch.Tensor) and b.is_nested):
for t1, t2 in zip(list(a), list(b), strict=True):
assert torch.equal(t1, t2), f"{msg} mismatch"
else:
assert torch.equal(a, b), f"{msg} mismatch"
class TestCheckpointRoundtrip:
"""Standard data → save → verify files → wipe → load → verify data."""
def test_tensor_fields(self, tq_system, checkpoint_dir, controller):
keys = ["k0", "k1"]
partition_id = "p_tensor"
input_ids = torch.tensor([[1, 2], [3, 4]])
attention_mask = torch.ones(2, 2)
tq.kv_batch_put(
keys=keys,
partition_id=partition_id,
fields=TensorDict({"input_ids": input_ids, "attention_mask": attention_mask}, batch_size=len(keys)),
tags=[{} for _ in keys],
)
tq.save_checkpoint(checkpoint_dir)
assert (checkpoint_dir / "metadata.json").exists()
assert (checkpoint_dir / "controller_state.pkl").exists()
su_dir = checkpoint_dir / "simple_storage"
assert su_dir.exists()
assert (su_dir / "storage_unit_info.json").exists()
with open(checkpoint_dir / "metadata.json") as f:
meta = json.load(f)
assert meta["storage_saved"] is True
ray.get(controller.clear_partition.remote(partition_id))
assert ray.get(controller.list_partitions.remote()) == []
tq.load_checkpoint(checkpoint_dir)
assert partition_id in ray.get(controller.list_partitions.remote())
retrieved = tq.kv_batch_get(keys=keys, partition_id=partition_id)
_assert_tensor_equal(retrieved["input_ids"], input_ids)
_assert_tensor_equal(retrieved["attention_mask"], attention_mask)
def test_controller_metadata(self, tq_system, checkpoint_dir, controller):
keys = ["a0", "a1", "a2"]
partition_id = "p_meta"
input_ids = torch.tensor([[1, 2, 3], [4, 5, 6], [7, 8, 9]])
tags = [{"idx": i} for i in range(3)]
tq.kv_batch_put(
keys=keys,
partition_id=partition_id,
fields=TensorDict({"input_ids": input_ids, "attention_mask": torch.ones(3, 3)}, batch_size=len(keys)),
tags=tags,
)
tq.save_checkpoint(checkpoint_dir)
ray.get(controller.clear_partition.remote(partition_id))
tq.load_checkpoint(checkpoint_dir)
snapshot = ray.get(controller.get_partition_snapshot.remote(partition_id))
for i, key in enumerate(keys):
assert key in snapshot.keys_mapping
gidx = snapshot.keys_mapping[key]
assert snapshot.custom_meta[gidx]["idx"] == i
def test_multiple_partitions(self, tq_system, checkpoint_dir, controller):
partitions_data = {f"part_{i}": (torch.full((2, 4), i, dtype=torch.long), torch.ones(2, 4)) for i in range(3)}
for pid, (iids, mask) in partitions_data.items():
tq.kv_batch_put(
keys=[f"{pid}_k0", f"{pid}_k1"],
partition_id=pid,
fields=TensorDict({"input_ids": iids, "attention_mask": mask}, batch_size=2),
tags=[{}, {}],
)
tq.save_checkpoint(checkpoint_dir)
for pid in partitions_data:
ray.get(controller.clear_partition.remote(pid))
tq.load_checkpoint(checkpoint_dir)
for pid, (iids, _) in partitions_data.items():
retrieved = tq.kv_batch_get(keys=[f"{pid}_k0", f"{pid}_k1"], partition_id=pid, select_fields=["input_ids"])
_assert_tensor_equal(retrieved["input_ids"], iids)
def test_user_metadata_preserved(self, tq_system, checkpoint_dir):
keys = ["m0"]
tq.kv_batch_put(
keys=keys,
partition_id="p_usermeta",
fields=TensorDict(
{"input_ids": torch.tensor([[10, 20]]), "attention_mask": torch.ones(1, 2)}, batch_size=1
),
tags=[{}],
)
tq.save_checkpoint(checkpoint_dir, metadata={"iteration": 42, "loss": 0.5})
with open(checkpoint_dir / "metadata.json") as f:
meta = json.load(f)
assert meta["user_metadata"]["iteration"] == 42
assert meta["user_metadata"]["loss"] == pytest.approx(0.5)
def test_non_tensor_fields(self, tq_system, checkpoint_dir, controller):
keys = ["t0", "t1"]
partition_id = "p_str"
input_ids = torch.tensor([[1, 2], [3, 4]])
fields = TensorDict(
{"input_ids": input_ids, "text": NonTensorStack("hello", "world")},
batch_size=2,
)
tq.kv_batch_put(keys=keys, partition_id=partition_id, fields=fields, tags=[{}, {}])
tq.save_checkpoint(checkpoint_dir)
ray.get(controller.clear_partition.remote(partition_id))
tq.load_checkpoint(checkpoint_dir)
retrieved = tq.kv_batch_get(keys=keys, partition_id=partition_id, select_fields=["input_ids"])
_assert_tensor_equal(retrieved["input_ids"], input_ids)
def test_nested_tensor_fields(self, tq_system, checkpoint_dir, controller):
keys = ["j0", "j1", "j2"]
partition_id = "p_jagged"
for i, key in enumerate(keys):
tq.kv_put(
key=key,
partition_id=partition_id,
fields=TensorDict({"seq": torch.arange(i + 1, dtype=torch.float).unsqueeze(0)}, batch_size=1),
tag=None,
)
tq.save_checkpoint(checkpoint_dir)
ray.get(controller.clear_partition.remote(partition_id))
tq.load_checkpoint(checkpoint_dir)
retrieved = tq.kv_batch_get(keys=keys, partition_id=partition_id, select_fields=["seq"])
for i, component in enumerate(retrieved["seq"].unbind()):
_assert_tensor_equal(component, torch.arange(i + 1, dtype=torch.float))
class TestIncludeStorageFalse:
"""For SimpleStorage, include_storage=False is silently forced to True."""
def test_storage_saved_is_true(self, tq_system, checkpoint_dir):
tq.kv_batch_put(
keys=["n0"],
partition_id="p_nometa",
fields=TensorDict({"input_ids": torch.tensor([[1, 2]]), "attention_mask": torch.ones(1, 2)}, batch_size=1),
tags=[{}],
)
tq.save_checkpoint(checkpoint_dir, include_storage=False)
with open(checkpoint_dir / "metadata.json") as f:
meta = json.load(f)
assert meta["storage_saved"] is True
assert (checkpoint_dir / "simple_storage").exists()
def test_both_restored_after_load(self, tq_system, checkpoint_dir, controller):
keys = ["n0", "n1"]
partition_id = "p_nometa2"
input_ids = torch.tensor([[5, 6], [7, 8]])
tq.kv_batch_put(
keys=keys,
partition_id=partition_id,
fields=TensorDict({"input_ids": input_ids, "attention_mask": torch.ones(2, 2)}, batch_size=len(keys)),
tags=[{} for _ in keys],
)
tq.save_checkpoint(checkpoint_dir, include_storage=False)
ray.get(controller.clear_partition.remote(partition_id))
tq.load_checkpoint(checkpoint_dir)
assert partition_id in ray.get(controller.list_partitions.remote())
snapshot = ray.get(controller.get_partition_snapshot.remote(partition_id))
for key in keys:
assert key in snapshot.keys_mapping
retrieved = tq.kv_batch_get(keys=keys, partition_id=partition_id)
_assert_tensor_equal(retrieved["input_ids"], input_ids)
class TestCheckpointErrors:
def test_save_raises_if_not_initialized(self, tmp_path):
import transfer_queue.interface as iface
original = iface._TQ_CONTROLLER
try:
iface._TQ_CONTROLLER = None
with pytest.raises(RuntimeError, match="not initialized"):
tq.save_checkpoint(tmp_path / "ck")
finally:
iface._TQ_CONTROLLER = original
def test_load_raises_if_not_initialized(self, tmp_path):
import transfer_queue.interface as iface
original = iface._TQ_CONTROLLER
try:
iface._TQ_CONTROLLER = None
with pytest.raises(RuntimeError, match="not initialized"):
tq.load_checkpoint(tmp_path / "ck")
finally:
iface._TQ_CONTROLLER = original
def test_load_raises_if_dir_missing(self, tq_system, tmp_path):
with pytest.raises(FileNotFoundError):
tq.load_checkpoint(tmp_path / "nonexistent")
def test_load_raises_if_metadata_missing(self, tq_system, tmp_path):
ck = tmp_path / "ck"
ck.mkdir()
with pytest.raises(FileNotFoundError, match="metadata.json"):
tq.load_checkpoint(ck)
def test_load_raises_on_storage_unit_count_mismatch(self, tq_system, tmp_path, checkpoint_dir):
tq.kv_batch_put(
keys=["e0"],
partition_id="p_err",
fields=TensorDict({"input_ids": torch.tensor([[1, 2]]), "attention_mask": torch.ones(1, 2)}, batch_size=1),
tags=[{}],
)
tq.save_checkpoint(checkpoint_dir)
su_info_path = checkpoint_dir / "simple_storage" / "storage_unit_info.json"
with open(su_info_path) as f:
su_info = json.load(f)
su_info.append({"position": 99, "storage_unit_id": "fake"})
with open(su_info_path, "w") as f:
json.dump(su_info, f)
with pytest.raises(ValueError, match="count mismatch"):
tq.load_checkpoint(checkpoint_dir)
def test_no_partial_state_on_failed_save(self, tq_system, tmp_path):
tq.kv_batch_put(
keys=["f0"],
partition_id="p_fail",
fields=TensorDict({"input_ids": torch.tensor([[1, 2]]), "attention_mask": torch.ones(1, 2)}, batch_size=1),
tags=[{}],
)
ck = tmp_path / "ck"
import unittest.mock as mock
with mock.patch(
"transfer_queue.client.TransferQueueClient.save_storage_checkpoint",
side_effect=RuntimeError("simulated dump failure"),
):
with pytest.raises(RuntimeError, match="simulated dump failure"):
tq.save_checkpoint(ck)
assert not ck.exists()
assert not (tmp_path / "ck.tmp").exists()