From 2ab66e112d596a61ec6b583ba7fc0d7911783abe Mon Sep 17 00:00:00 2001 From: Evan Date: Tue, 3 Mar 2026 15:25:57 +0000 Subject: [PATCH] or maybe this --- src/exo/master/placement.py | 12 +++++++++++- src/exo/master/placement_utils.py | 7 ++++++- src/exo/shared/types/topology.py | 1 + 3 files changed, 18 insertions(+), 2 deletions(-) diff --git a/src/exo/master/placement.py b/src/exo/master/placement.py index d6f382ec..ba058903 100644 --- a/src/exo/master/placement.py +++ b/src/exo/master/placement.py @@ -69,7 +69,17 @@ def place_instance( required_nodes: set[NodeId] | None = None, ) -> dict[InstanceId, Instance]: cycles = topology.get_cycles() - candidate_cycles = list(filter(lambda it: len(it) >= command.min_nodes, cycles)) + candidate_cycles = filter(lambda it: len(it) >= command.min_nodes, cycles) + nodes_with_exo = set[NodeId]() + for instance in current_instances.values(): + nodes_with_exo |= instance.shard_assignments.node_to_runner.keys() + candidate_cycles = [ + Cycle( + cycle.node_ids, + has_other_placement=any(node in nodes_with_exo for node in cycle), + ) + for cycle in candidate_cycles + ] # Filter to cycles containing all required nodes (subset matching) if required_nodes: diff --git a/src/exo/master/placement_utils.py b/src/exo/master/placement_utils.py index a8a2e70b..ff800c5e 100644 --- a/src/exo/master/placement_utils.py +++ b/src/exo/master/placement_utils.py @@ -29,7 +29,12 @@ def filter_cycles_by_memory( continue total_mem = sum( - (node_memory[node_id].ram_total for node_id in cycle.node_ids), + ( + node_memory[node_id].ram_available + if cycle.has_other_placement + else node_memory[node_id].ram_total + for node_id in cycle.node_ids + ), start=Memory(), ) if total_mem >= required_memory: diff --git a/src/exo/shared/types/topology.py b/src/exo/shared/types/topology.py index c388c6ab..e2f29d02 100644 --- a/src/exo/shared/types/topology.py +++ b/src/exo/shared/types/topology.py @@ -9,6 +9,7 @@ from exo.utils.pydantic_ext import FrozenModel @dataclass(frozen=True) class Cycle: node_ids: list[NodeId] + has_other_placement: bool = False def __len__(self) -> int: return self.node_ids.__len__()