[object Object]

← back to Exo

Fix placement filter to use subset matching instead of exact match (#1265)

d229df38f93f7d9f4e255abfec895a7e72b51aac · 2026-01-23 19:40:31 +0000 · Alex Cheema

## Motivation

When using the dashboard's instance placement filter (clicking nodes in
the topology), it was filtering to placements that use exactly the
selected nodes. This isn't the expected behavior - users want to see
placements that include all selected nodes, but may also include
additional nodes.

For example, selecting nodes [A, B] should show placements using [A, B],
[A, B, C], [A, B, C, D], etc. - not just [A, B].

## Changes

- Added `required_nodes` parameter to `place_instance()` in
`placement.py`
- Filter cycles early in placement to only those containing all required
nodes (subset matching)
- Simplified `api.py` by removing the subgraph topology filtering and
passing `required_nodes` directly to placement
- Renamed internal `node_ids` variable to `placement_node_ids` to avoid
shadowing the parameter

## Why It Works

By filtering cycles at the placement level using
`required_nodes.issubset(cycle.node_ids)`, we ensure that only cycles
containing all the user-selected nodes are considered. This happens
early in the placement algorithm, so we don't waste time computing
placements that would be filtered out later.

## Test Plan

### Manual Testing
- Select nodes in the dashboard topology view
- Verify that placements shown include all selected nodes (but may
include additional nodes)
- Verify that placements not containing the selected nodes are filtered
out

### Automated Testing
- Existing placement tests pass
- `uv run pytest src/exo/master/tests/ -v` - 37 tests pass

Co-authored-by: Claude Opus 4.5 <noreply@anthropic.com>

Files touched

Diff

commit d229df38f93f7d9f4e255abfec895a7e72b51aac
Author: Alex Cheema <41707476+AlexCheema@users.noreply.github.com>
Date:   Fri Jan 23 19:40:31 2026 +0000

    Fix placement filter to use subset matching instead of exact match (#1265)
    
    ## Motivation
    
    When using the dashboard's instance placement filter (clicking nodes in
    the topology), it was filtering to placements that use exactly the
    selected nodes. This isn't the expected behavior - users want to see
    placements that include all selected nodes, but may also include
    additional nodes.
    
    For example, selecting nodes [A, B] should show placements using [A, B],
    [A, B, C], [A, B, C, D], etc. - not just [A, B].
    
    ## Changes
    
    - Added `required_nodes` parameter to `place_instance()` in
    `placement.py`
    - Filter cycles early in placement to only those containing all required
    nodes (subset matching)
    - Simplified `api.py` by removing the subgraph topology filtering and
    passing `required_nodes` directly to placement
    - Renamed internal `node_ids` variable to `placement_node_ids` to avoid
    shadowing the parameter
    
    ## Why It Works
    
    By filtering cycles at the placement level using
    `required_nodes.issubset(cycle.node_ids)`, we ensure that only cycles
    containing all the user-selected nodes are considered. This happens
    early in the placement algorithm, so we don't waste time computing
    placements that would be filtered out later.
    
    ## Test Plan
    
    ### Manual Testing
    - Select nodes in the dashboard topology view
    - Verify that placements shown include all selected nodes (but may
    include additional nodes)
    - Verify that placements not containing the selected nodes are filtered
    out
    
    ### Automated Testing
    - Existing placement tests pass
    - `uv run pytest src/exo/master/tests/ -v` - 37 tests pass
    
    Co-authored-by: Claude Opus 4.5 <noreply@anthropic.com>
---
 src/exo/master/api.py       | 39 +++++++++++++++++++++++----------------
 src/exo/master/placement.py |  9 +++++++++
 2 files changed, 32 insertions(+), 16 deletions(-)

diff --git a/src/exo/master/api.py b/src/exo/master/api.py
index e59c293e..3d56887f 100644
--- a/src/exo/master/api.py
+++ b/src/exo/master/api.py
@@ -356,14 +356,9 @@ class API:
     ) -> PlacementPreviewResponse:
         seen: set[tuple[ModelId, Sharding, InstanceMeta, int]] = set()
         previews: list[PlacementPreview] = []
+        required_nodes = set(node_ids) if node_ids else None
 
-        # Create filtered topology if node_ids specified
-        if node_ids and len(node_ids) > 0:
-            topology = self.state.topology.get_subgraph_from_nodes(node_ids)
-        else:
-            topology = self.state.topology
-
-        if len(list(topology.list_nodes())) == 0:
+        if len(list(self.state.topology.list_nodes())) == 0:
             return PlacementPreviewResponse(previews=[])
 
         cards = [card for card in MODEL_CARDS.values() if card.model_id == model_id]
@@ -376,7 +371,9 @@ class API:
                 instance_combinations.extend(
                     [
                         (sharding, instance_meta, i)
-                        for i in range(1, len(list(topology.list_nodes())) + 1)
+                        for i in range(
+                            1, len(list(self.state.topology.list_nodes())) + 1
+                        )
                     ]
                 )
         # TODO: PDD
@@ -394,8 +391,9 @@ class API:
                         ),
                         node_memory=self.state.node_memory,
                         node_network=self.state.node_network,
-                        topology=topology,
+                        topology=self.state.topology,
                         current_instances=self.state.instances,
+                        required_nodes=required_nodes,
                     )
                 except ValueError as exc:
                     if (model_card.model_id, sharding, instance_meta, 0) not in seen:
@@ -434,14 +432,16 @@ class API:
 
                 instance = new_instances[0]
                 shard_assignments = instance.shard_assignments
-                node_ids = list(shard_assignments.node_to_runner.keys())
+                placement_node_ids = list(shard_assignments.node_to_runner.keys())
 
                 memory_delta_by_node: dict[str, int] = {}
-                if node_ids:
+                if placement_node_ids:
                     total_bytes = model_card.storage_size.in_bytes
-                    per_node = total_bytes // len(node_ids)
-                    remainder = total_bytes % len(node_ids)
-                    for index, node_id in enumerate(sorted(node_ids, key=str)):
+                    per_node = total_bytes // len(placement_node_ids)
+                    remainder = total_bytes % len(placement_node_ids)
+                    for index, node_id in enumerate(
+                        sorted(placement_node_ids, key=str)
+                    ):
                         extra = 1 if index < remainder else 0
                         memory_delta_by_node[str(node_id)] = per_node + extra
 
@@ -449,7 +449,7 @@ class API:
                     model_card.model_id,
                     sharding,
                     instance_meta,
-                    len(node_ids),
+                    len(placement_node_ids),
                 ) not in seen:
                     previews.append(
                         PlacementPreview(
@@ -461,7 +461,14 @@ class API:
                             error=None,
                         )
                     )
-                seen.add((model_card.model_id, sharding, instance_meta, len(node_ids)))
+                seen.add(
+                    (
+                        model_card.model_id,
+                        sharding,
+                        instance_meta,
+                        len(placement_node_ids),
+                    )
+                )
 
         return PlacementPreviewResponse(previews=previews)
 
diff --git a/src/exo/master/placement.py b/src/exo/master/placement.py
index c8050c6e..e37f533f 100644
--- a/src/exo/master/placement.py
+++ b/src/exo/master/placement.py
@@ -54,9 +54,18 @@ def place_instance(
     current_instances: Mapping[InstanceId, Instance],
     node_memory: Mapping[NodeId, MemoryUsage],
     node_network: Mapping[NodeId, NodeNetworkInfo],
+    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))
+
+    # Filter to cycles containing all required nodes (subset matching)
+    if required_nodes:
+        candidate_cycles = [
+            cycle
+            for cycle in candidate_cycles
+            if required_nodes.issubset(cycle.node_ids)
+        ]
     cycles_with_sufficient_memory = filter_cycles_by_memory(
         candidate_cycles, node_memory, command.model_card.storage_size
     )

← 8a595fee Fix Thunderbolt bridge cycle detection to include 2-node cyc  ·  back to Exo  ·  Add FLUX.1-Krea-dev model (#1269) 23fd37fe →