diff --git a/src/exo/master/placement.py b/src/exo/master/placement.py index 0ed634e2ee..c5dc063f09 100644 --- a/src/exo/master/placement.py +++ b/src/exo/master/placement.py @@ -103,6 +103,25 @@ def _cycle_download_score( ) +def _not_enough_memory( + command: PlaceInstance, + cycles: Sequence[Cycle], + node_memory: Mapping[NodeId, MemoryUsage], +) -> str: + reported = [cycle for cycle in cycles if all(node in node_memory for node in cycle)] + if not reported: + return "Waiting for the nodes to report their memory" + most_free = max( + sum((node_memory[node].ram_available for node in cycle), start=Memory()) + for cycle in reported + ) + return ( + f"Not enough memory: {command.model_card.model_id} needs " + f"{command.model_card.storage_size.in_gb:.1f} GB, but connected nodes have at most " + f"{most_free.in_gb:.1f} GB free between them" + ) + + def place_instance( command: PlaceInstance, topology: Topology, @@ -115,7 +134,14 @@ def place_instance( node_rdma_ctl: Mapping[NodeId, NodeRdmaCtlStatus] | None = None, ) -> dict[InstanceId, Instance]: cycles = topology.get_cycles() + if not cycles: + raise ValueError("No nodes are available") candidate_cycles = list(filter(lambda it: len(it) >= command.min_nodes, cycles)) + if not candidate_cycles: + largest = max(len(cycle) for cycle in cycles) + raise ValueError( + f"Needs {command.min_nodes} connected nodes, but at most {largest} are connected to each other" + ) # Filter to cycles containing all required nodes (subset matching) if required_nodes: @@ -124,11 +150,13 @@ def place_instance( for cycle in candidate_cycles if required_nodes.issubset(cycle.node_ids) ] + if not candidate_cycles: + raise ValueError("The chosen nodes aren't all connected to each other") cycles_with_sufficient_memory = filter_cycles_by_memory( candidate_cycles, node_memory, command.model_card.storage_size ) if len(cycles_with_sufficient_memory) == 0: - raise ValueError("No cycles found with sufficient memory") + raise ValueError(_not_enough_memory(command, candidate_cycles, node_memory)) if command.sharding == Sharding.Tensor: if not command.model_card.supports_tensor: diff --git a/src/exo/master/tests/test_placement.py b/src/exo/master/tests/test_placement.py index b891187266..287390af3d 100644 --- a/src/exo/master/tests/test_placement.py +++ b/src/exo/master/tests/test_placement.py @@ -271,10 +271,143 @@ def test_get_instance_placements_one_node_not_fit() -> None: ), ) - with pytest.raises(ValueError, match="No cycles found with sufficient memory"): + with pytest.raises(ValueError) as error: place_instance( cic, topology, {}, node_memory, node_network, _metal_only(node_memory) ) + assert str(error.value) == ( + "Not enough memory: test-model needs 0.0 GB, but connected nodes have at most " + "0.0 GB free between them" + ) + + +def _connected(*groups: list[NodeId]) -> Topology: + """Nodes connected both ways to every other node in their group.""" + topology = Topology() + port = 0 + for group in groups: + for node_id in group: + topology.add_node(node_id) + for source in group: + for sink in group: + if source != sink: + port += 1 + topology.add_connection( + Connection( + source=source, + sink=sink, + edge=create_socket_connection(port), + ) + ) + return topology + + +def _small_model() -> ModelCard: + return ModelCard( + model_id=ModelId("test-model"), + storage_size=Memory.from_kb(1000), + n_layers=10, + hidden_size=1000, + supports_tensor=True, + tasks=[ModelTask.TextGeneration], + backends=[Backend.MlxMetal], + ) + + +def test_placement_says_when_too_few_nodes_are_connected() -> None: + a, b, c = NodeId(), NodeId(), NodeId() + topology = _connected([a, b], [c]) + node_memory = {n: create_node_memory(10**9) for n in (a, b, c)} + node_network = {n: create_node_network() for n in (a, b, c)} + command = place_instance_command(_small_model()).model_copy(update={"min_nodes": 3}) + + with pytest.raises(ValueError) as error: + place_instance( + command, + topology, + {}, + node_memory, + node_network, + _metal_only(node_memory), + ) + + assert str(error.value) == ( + "Needs 3 connected nodes, but at most 2 are connected to each other" + ) + + +def test_placement_says_when_there_are_no_nodes() -> None: + with pytest.raises(ValueError) as error: + place_instance( + place_instance_command(_small_model()), Topology(), {}, {}, {}, {} + ) + + assert str(error.value) == "No nodes are available" + + +def test_placement_says_when_the_chosen_nodes_are_not_connected() -> None: + a, b, c = NodeId(), NodeId(), NodeId() + topology = _connected([a, b], [c]) + node_memory = {n: create_node_memory(10**9) for n in (a, b, c)} + node_network = {n: create_node_network() for n in (a, b, c)} + + with pytest.raises(ValueError) as error: + place_instance( + place_instance_command(_small_model()), + topology, + {}, + node_memory, + node_network, + _metal_only(node_memory), + required_nodes={a, c}, + ) + + assert str(error.value) == ("The chosen nodes aren't all connected to each other") + + +def test_placement_says_how_much_memory_is_free() -> None: + a, b = NodeId(), NodeId() + topology = _connected([a, b]) + node_memory = { + a: create_node_memory(300 * 1024**2), + b: create_node_memory(500 * 1024**2), + } + node_network = {n: create_node_network() for n in (a, b)} + model = _small_model().model_copy( + update={"storage_size": Memory.from_bytes(2 * 1024**3)} + ) + + with pytest.raises(ValueError) as error: + place_instance( + place_instance_command(model), + topology, + {}, + node_memory, + node_network, + _metal_only(node_memory), + ) + + assert str(error.value) == ( + "Not enough memory: test-model needs 2.0 GB, but connected nodes have at most " + "0.8 GB free between them" + ) + + +def test_placement_says_when_nodes_have_not_reported_their_memory() -> None: + a, b = NodeId(), NodeId() + topology = _connected([a, b]) + + with pytest.raises(ValueError) as error: + place_instance( + place_instance_command(_small_model()), + topology, + {}, + {}, + {}, + {a: [Backend.MlxMetal], b: [Backend.MlxMetal]}, + ) + + assert str(error.value) == "Waiting for the nodes to report their memory" def test_get_transition_events_no_change(instance: Instance):