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
30 changes: 29 additions & 1 deletion src/exo/master/placement.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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:
Expand All @@ -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:
Expand Down
135 changes: 134 additions & 1 deletion src/exo/master/tests/test_placement.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
Loading