diff --git a/src/zeroband/utils/state_dict_send_recv.py b/src/zeroband/utils/state_dict_send_recv.py index 66366dd9..5ca3f19d 100644 --- a/src/zeroband/utils/state_dict_send_recv.py +++ b/src/zeroband/utils/state_dict_send_recv.py @@ -152,12 +152,13 @@ def recv_state_dict(pg: ProcessGroup, src_rank: int, og_state_dict: dict) -> dic for job in jobs: job.wait() + # Clone received payloads so the returned state dict does not alias og_state_dict tensors. + new_tensors: list[torch.Tensor] = [] for tensor, data in zip(tensors, datas): - if isinstance(tensor, DTensor): - tensor = tensor.to_local() - tensor.copy_(data) + buffer = tensor.to_local() if isinstance(tensor, DTensor) else tensor + new_tensors.append(data.clone().to(dtype=buffer.dtype)) - state_dict = _load_sendable_state_dict(tensors, state_dict) + state_dict = _load_sendable_state_dict(new_tensors, state_dict) # logger = get_logger() # logger.debug(f"recv tensors {get_tensor_list_signature(tensors)}") diff --git a/tests/test_dist/test_send_state_dict.py b/tests/test_dist/test_send_state_dict.py index e4e1f22f..0698baeb 100644 --- a/tests/test_dist/test_send_state_dict.py +++ b/tests/test_dist/test_send_state_dict.py @@ -46,6 +46,35 @@ def test_load_state_dict(): assert state_dict_to_send["nested_data"]["foo"] == state_dict_copy["nested_data"]["foo"] +def test_recv_state_dict_no_tensor_aliasing(): + """recv_state_dict must not return tensors that alias og_state_dict buffers.""" + og_state_dict = { + "step": 10, + "world": "template", + "optim_sates": torch.zeros(10), + "nested_data": {"foo": "bar", "tensor": torch.zeros(10)}, + } + received_payload = { + "step": 0, + "world": "karl is having his best life", + "optim_sates": torch.ones(10), + "nested_data": {"foo": "bar", "tensor": torch.ones(10)}, + } + + non_tensored_state, template_tensors = _get_sendable_state_dict(og_state_dict) + _, payload_tensors = _get_sendable_state_dict(received_payload) + + new_tensors = [data.clone() for data in payload_tensors] + result = _load_sendable_state_dict(new_tensors, non_tensored_state) + + assert (result["optim_sates"] == received_payload["optim_sates"]).all() + assert id(result["optim_sates"]) != id(og_state_dict["optim_sates"]) + assert (result["nested_data"]["tensor"] == received_payload["nested_data"]["tensor"]).all() + assert id(result["nested_data"]["tensor"]) != id(og_state_dict["nested_data"]["tensor"]) + assert result["step"] == received_payload["step"] + assert result["world"] == received_payload["world"] + + @pytest.mark.skip(reason="hang") @pytest.mark.parametrize("world_size", [2]) def test_send_recv_state_dict(world_size: int, random_available_port: int, mock_env): @@ -70,19 +99,19 @@ def foo(**kwargs): rank = int(os.environ.get("RANK")) if rank == 0: - send_state_dict(state_dict_to_send, 1, world_size) + send_state_dict(edm.global_pg, state_dict_to_send, 1) else: - state_dict = recv_state_dict(pg=edm.global_pg, rank=0, world_size=world_size) + state_dict = recv_state_dict(edm.global_pg, 0, state_dict_to_recv) - assert (state_dict["optim_sates"] == state_dict_to_recv["optim_sates"]).all() + assert (state_dict["optim_sates"] == state_dict_to_send["optim_sates"]).all() assert id(state_dict["optim_sates"]) != id(state_dict_to_recv["optim_sates"]) - assert (state_dict["nested_data"]["tensor"] == state_dict_to_recv["nested_data"]["tensor"]).all() + assert (state_dict["nested_data"]["tensor"] == state_dict_to_send["nested_data"]["tensor"]).all() assert id(state_dict["nested_data"]["tensor"]) != id(state_dict_to_recv["nested_data"]["tensor"]) - assert state_dict["step"] == state_dict_to_recv["step"] - assert state_dict["world"] == state_dict_to_recv["world"] - assert state_dict["nested_data"]["foo"] == state_dict_to_recv["nested_data"]["foo"] + assert state_dict["step"] == state_dict_to_send["step"] + assert state_dict["world"] == state_dict_to_send["world"] + assert state_dict["nested_data"]["foo"] == state_dict_to_send["nested_data"]["foo"] del edm