Skip to content
Open
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
9 changes: 5 additions & 4 deletions src/zeroband/utils/state_dict_send_recv.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)}")
Expand Down
43 changes: 36 additions & 7 deletions tests/test_dist/test_send_state_dict.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand All @@ -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

Expand Down