[object Object]

← back to Exo

Worker refactor

aa519b8c0339b33f8f41f797459b3295a8b2891c · 2025-11-10 23:31:53 +0000 · Evan Quiney

Co-authored-by: rltakashige <rl.takashige@gmail.com>
Co-authored-by: Alex Cheema <alexcheema123@gmail.com>

Files touched

Diff

commit aa519b8c0339b33f8f41f797459b3295a8b2891c
Author: Evan Quiney <evanev7@gmail.com>
Date:   Mon Nov 10 23:31:53 2025 +0000

    Worker refactor
    
    Co-authored-by: rltakashige <rl.takashige@gmail.com>
    Co-authored-by: Alex Cheema <alexcheema123@gmail.com>
---
 .github/configs/bench_simple.yaml                  |  33 +-
 .github/scripts/bench.py                           |  90 +++-
 .mlx_typings/mlx_lm/__init__.pyi                   |   1 +
 .mlx_typings/mlx_lm/generate.pyi                   |   2 +-
 dashboard/index.html                               | 186 ++++---
 flake.nix                                          |   3 +-
 pyproject.toml                                     |   2 +-
 src/exo/engines/mlx/auto_parallel.py               | 292 +++++-----
 src/exo/engines/mlx/utils_mlx.py                   | 197 ++++---
 src/exo/main.py                                    |  14 +
 src/exo/master/api.py                              |  66 ++-
 src/exo/master/main.py                             |  10 +-
 src/exo/master/placement.py                        |  72 +--
 src/exo/master/placement_utils.py                  |  44 +-
 src/exo/shared/apply.py                            |  90 ++--
 src/exo/shared/global_conn.py                      |  67 ---
 src/exo/shared/logging.py                          |  21 +-
 src/exo/shared/types/api.py                        |   8 +-
 src/exo/shared/types/chunks.py                     |   2 -
 src/exo/shared/types/commands.py                   |  10 +-
 src/exo/shared/types/events.py                     |  78 +--
 src/exo/shared/types/memory.py                     |   5 +
 src/exo/shared/types/state.py                      |  13 +-
 src/exo/shared/types/tasks.py                      |  39 +-
 src/exo/shared/types/worker/commands_runner.py     |  26 -
 src/exo/shared/types/worker/common.py              |  25 -
 src/exo/shared/types/worker/communication.py       |  40 --
 src/exo/shared/types/worker/downloads.py           |   2 +
 src/exo/shared/types/worker/instances.py           |  62 ++-
 src/exo/shared/types/worker/ops.py                 |  39 +-
 .../types/worker/parallelisation_strategy.py       |  13 -
 src/exo/shared/types/worker/runners.py             |  44 +-
 src/exo/shared/types/worker/shards.py              |  25 +-
 src/exo/utils/channels.py                          | 224 +++++++-
 src/exo/utils/pydantic_ext.py                      |   4 +
 src/exo/utils/tests/testing_mp.py                  |  41 ++
 src/exo/worker/common.py                           |  36 --
 src/exo/worker/download/shard_downloader.py        |   1 -
 src/exo/worker/main.py                             | 588 ++++++---------------
 src/exo/worker/plan.py                             | 436 ++++++---------
 src/exo/worker/runner/bootstrap.py                 |  33 +-
 src/exo/worker/runner/generate.py                  | 411 ++------------
 src/exo/worker/runner/runner.py                    | 293 ++++++----
 src/exo/worker/runner/runner_supervisor.py         | 386 +++++---------
 src/exo/worker/runner/utils.py                     |   4 +-
 tmp/run_llm.sh                                     |   4 +-
 uv.lock                                            |   2 +-
 47 files changed, 1751 insertions(+), 2333 deletions(-)

diff --git a/.github/configs/bench_simple.yaml b/.github/configs/bench_simple.yaml
index 7ba45e44..346df681 100644
--- a/.github/configs/bench_simple.yaml
+++ b/.github/configs/bench_simple.yaml
@@ -10,7 +10,7 @@ hardware_plan:
 environment:
   PLACEHOLDER: "placeholder"
   # OVERRIDE_MEMORY_MB: 50000
-  # MLX_METAL_FAST_SYNCH: 1
+  MLX_METAL_FAST_SYNCH: 1
 
 # Timeout for instance and runner readiness (seconds)
 timeout_seconds: 1800
@@ -18,12 +18,17 @@ timeout_seconds: 1800
 # Model instances to run concurrently
 model_ids:
   # - "mlx-community/DeepSeek-V3.1-8bit"
+  - "mlx-community/Kimi-K2-Instruct-4bit"
   # - "mlx-community/Qwen3-235B-A22B-4bit"
   # - "mlx-community/Llama-3.3-70B-Instruct-4bit"
-  - "mlx-community/Llama-3.3-70B-Instruct-8bit"
+  # - "mlx-community/Llama-3.3-70B-Instruct-8bit"
+  # - "mlx-community/Llama-3.2-1B-Instruct-4bit"
 
-# Placement strategy: "tensor", "pipeline", or "auto"
-strategy: "pipeline"
+# Sharding strategy: "Pipeline" or "Tensor"
+sharding: "Tensor"
+
+# Instance type: "MlxRing" or "MlxIbv"
+instance_meta: "MlxIbv"
 
 # If true, run requests sequentially (no overlap); if false, fire-and-forget (default: false)
 no_overlap: true
@@ -97,13 +102,13 @@ stages:
     generation_length: 512
     time_between_requests: 2.0
     iterations: 10
-  - name: "pp16384_g64"
-    prompt_length: 16384
-    generation_length: 64
-    time_between_requests: 2.0
-    iterations: 10
-  - name: "pp16384_g512"
-    prompt_length: 16384
-    generation_length: 512
-    time_between_requests: 2.0
-    iterations: 10
+  # - name: "pp16384_g64"
+  #   prompt_length: 16384
+  #   generation_length: 64
+  #   time_between_requests: 2.0
+  #   iterations: 10
+  # - name: "pp16384_g512"
+  #   prompt_length: 16384
+  #   generation_length: 512
+  #   time_between_requests: 2.0
+  #   iterations: 10
diff --git a/.github/scripts/bench.py b/.github/scripts/bench.py
index 6a841fbb..06b81542 100644
--- a/.github/scripts/bench.py
+++ b/.github/scripts/bench.py
@@ -100,6 +100,23 @@ def fetch_state(api_base: str) -> dict[str, Any]:
     return _http_request(f"{api_base}/state")
 
 
+def unwrap_tagged_union(obj: Any) -> tuple[str | None, Any]:
+    """Extract tag and payload from tagged union format {Tag: {fields...}}.
+    
+    Returns (tag_name, payload) if the object is a tagged union, otherwise (None, obj).
+    """
+    if not isinstance(obj, dict):
+        return None, obj
+    
+    keys = list(obj.keys())
+    if len(keys) == 1 and isinstance(keys[0], str):
+        tag = keys[0]
+        payload = obj[tag]
+        return tag, payload
+    
+    return None, obj
+
+
 def collect_metrics_snapshot(state: Mapping[str, Any]) -> MetricsSnapshot:
     """Collect current metrics snapshot from state."""
     timestamp = time.time()
@@ -144,7 +161,9 @@ def collect_metrics_snapshot(state: Mapping[str, Any]) -> MetricsSnapshot:
     
     # Map instance_id -> node_ids (instances can span multiple nodes)
     instance_to_nodes: dict[str, set[str]] = {}
-    for instance_id, instance_data in instances.items():
+    for instance_id, instance_wrapped in instances.items():
+        # Unwrap tagged Instance union (MlxRingInstance or MlxIbvInstance)
+        _instance_tag, instance_data = unwrap_tagged_union(instance_wrapped)
         if not isinstance(instance_data, dict):
             continue
         
@@ -175,7 +194,7 @@ def collect_metrics_snapshot(state: Mapping[str, Any]) -> MetricsSnapshot:
             tasks_skipped += 1
             continue
         
-        # Extract actual task from wrapper (e.g., {"ChatCompletionTask": {...}})
+        # Extract actual task from wrapper (e.g., {"ChatCompletion": {...}})
         if len(task_wrapper) != 1:
             print(f"[DEBUG] Task wrapper has unexpected number of keys: {len(task_wrapper)}")
             tasks_skipped += 1
@@ -279,9 +298,14 @@ def count_instances_by_model(state: Mapping[str, Any], model_id: str) -> int:
     """Count how many instances exist for a given model_id."""
     instances: Mapping[str, Any] = state.get("instances", {})
     count = 0
-    for instance in instances.values():
-        shard = instance.get("shardAssignments", {})
-        if shard.get("modelId") == model_id:
+    for instance_wrapped in instances.values():
+        # Unwrap tagged Instance union
+        _instance_tag, instance_data = unwrap_tagged_union(instance_wrapped)
+        if not isinstance(instance_data, dict):
+            continue
+        
+        shard = instance_data.get("shardAssignments", {})
+        if isinstance(shard, dict) and shard.get("modelId") == model_id:
             count += 1
     return count
 
@@ -290,9 +314,14 @@ def get_all_instance_ids_for_model(state: Mapping[str, Any], model_id: str) -> l
     """Get all instance IDs for a given model_id."""
     instances: Mapping[str, Any] = state.get("instances", {})
     instance_ids = []
-    for instance_id, instance in instances.items():
-        shard = instance.get("shardAssignments", {})
-        if shard.get("modelId") == model_id:
+    for instance_id, instance_wrapped in instances.items():
+        # Unwrap tagged Instance union
+        _instance_tag, instance_data = unwrap_tagged_union(instance_wrapped)
+        if not isinstance(instance_data, dict):
+            continue
+        
+        shard = instance_data.get("shardAssignments", {})
+        if isinstance(shard, dict) and shard.get("modelId") == model_id:
             instance_ids.append(instance_id)
     return instance_ids
 
@@ -302,9 +331,14 @@ def count_ready_instances_by_model(state: Mapping[str, Any], model_id: str) -> i
     instances: Mapping[str, Any] = state.get("instances", {})
     ready_count = 0
     
-    for instance_id, instance in instances.items():
-        shard = instance.get("shardAssignments", {})
-        if shard.get("modelId") != model_id:
+    for instance_id, instance_wrapped in instances.items():
+        # Unwrap tagged Instance union
+        _instance_tag, instance_data = unwrap_tagged_union(instance_wrapped)
+        if not isinstance(instance_data, dict):
+            continue
+        
+        shard = instance_data.get("shardAssignments", {})
+        if not isinstance(shard, dict) or shard.get("modelId") != model_id:
             continue
         
         # Check if all runners for this instance are ready
@@ -312,8 +346,9 @@ def count_ready_instances_by_model(state: Mapping[str, Any], model_id: str) -> i
         if len(runner_ids) == 0:
             continue
         
+        # Fixed runner status names: RunnerReady and RunnerRunning (not LoadedRunnerStatus/RunningRunnerStatus)
         all_ready = all(
-            get_runner_status_kind(state, rid) in {"LoadedRunnerStatus", "RunningRunnerStatus"}
+            get_runner_status_kind(state, rid) in {"RunnerReady", "RunnerRunning"}
             for rid in runner_ids
         )
         
@@ -325,8 +360,18 @@ def count_ready_instances_by_model(state: Mapping[str, Any], model_id: str) -> i
 
 def get_runner_ids_for_instance(state: Mapping[str, Any], instance_id: str) -> list[str]:
     instances: Mapping[str, Any] = state.get("instances", {})
-    inst = instances.get(instance_id, {})
-    r2s = inst.get("shardAssignments", {}).get("runnerToShard", {})
+    instance_wrapped = instances.get(instance_id, {})
+    
+    # Unwrap tagged Instance union
+    _instance_tag, instance_data = unwrap_tagged_union(instance_wrapped)
+    if not isinstance(instance_data, dict):
+        return []
+    
+    shard_assignments = instance_data.get("shardAssignments", {})
+    if not isinstance(shard_assignments, dict):
+        return []
+    
+    r2s = shard_assignments.get("runnerToShard", {})
     if isinstance(r2s, dict):
         return list(r2s.keys())
     return []
@@ -860,8 +905,9 @@ async def run_benchmark(
     else:
         raise ValueError("Config must contain either 'model_id' or 'model_ids'")
     
-    # Get strategy (optional, defaults to None if not specified)
-    strategy: str | None = config.get("strategy")
+    # Get sharding and instance_meta (optional, defaults to None if not specified)
+    sharding: str | None = config.get("sharding")
+    instance_meta: str | None = config.get("instance_meta")
     
     # Get no_overlap flag (optional, defaults to False)
     no_overlap: bool = config.get("no_overlap", False)
@@ -874,7 +920,8 @@ async def run_benchmark(
     print(f"Configuration File: {config_path}")
     print(f"Model IDs:          {model_ids}")
     print(f"Instance Count:     {len(model_ids)}")
-    print(f"Strategy:           {strategy if strategy else 'not specified'}")
+    print(f"Sharding:           {sharding if sharding else 'not specified (defaults to Pipeline)'}")
+    print(f"Instance Type:      {instance_meta if instance_meta else 'not specified (defaults to MlxRing)'}")
     print(f"No Overlap:         {no_overlap}")
     print(f"Stages:             {len(stages)}")
     print(f"Expected Nodes:     {expected_nodes}")
@@ -916,8 +963,10 @@ async def run_benchmark(
                 
                 # Build instance creation request data
                 instance_data: dict[str, Any] = {"model_id": model_id}
-                if strategy is not None:
-                    instance_data["strategy"] = strategy
+                if sharding is not None:
+                    instance_data["sharding"] = sharding
+                if instance_meta is not None:
+                    instance_data["instance_meta"] = instance_meta
                 
                 response = await _http_request_async(
                     f"{api_base}/instance",
@@ -1027,7 +1076,8 @@ async def run_benchmark(
                     "instance_ids": all_instance_ids,
                     "instance_count": len(all_instance_ids),
                     "runner_count": total_runners,
-                    "strategy": strategy,
+                    "sharding": sharding,
+                    "instance_meta": instance_meta,
                 },
                 "configuration": {
                     "stages": [
diff --git a/.mlx_typings/mlx_lm/__init__.pyi b/.mlx_typings/mlx_lm/__init__.pyi
index fee89807..2ed43899 100644
--- a/.mlx_typings/mlx_lm/__init__.pyi
+++ b/.mlx_typings/mlx_lm/__init__.pyi
@@ -1,2 +1,3 @@
 import models as models
 import tokenizer_utils as tokenizer_utils
+from generate import *
diff --git a/.mlx_typings/mlx_lm/generate.pyi b/.mlx_typings/mlx_lm/generate.pyi
index 4711fce0..8a957608 100644
--- a/.mlx_typings/mlx_lm/generate.pyi
+++ b/.mlx_typings/mlx_lm/generate.pyi
@@ -175,7 +175,7 @@ def stream_generate(
     prompt: Union[str, mx.array, List[int]],
     max_tokens: int = ...,
     draft_model: Optional[nn.Module] = ...,
-    **kwargs,
+    **kwargs: object,
 ) -> Generator[GenerationResponse, None, None]:
     """
     A generator producing text based on the given prompt from the model.
diff --git a/dashboard/index.html b/dashboard/index.html
index 715fdb54..62ec32f5 100644
--- a/dashboard/index.html
+++ b/dashboard/index.html
@@ -969,27 +969,29 @@
                 </select>
                 
                 <div class="strategy-selector">
-                    <label class="launch-label">Parallelization Strategy:</label>
+                    <label class="launch-label">Sharding:</label>
                     <div class="strategy-options">
                         <div class="strategy-option">
-                            <input type="radio" id="strategyAuto" name="strategy" value="auto" checked>
-                            <label for="strategyAuto">Auto</label>
+                            <input type="radio" id="shardingPipeline" name="sharding" value="Pipeline" checked>
+                            <label for="shardingPipeline">Pipeline</label>
                         </div>
                         <div class="strategy-option">
-                            <input type="radio" id="strategyPipeline" name="strategy" value="pipeline">
-                            <label for="strategyPipeline">Pipeline</label>
-                        </div>
-                        <div class="strategy-option">
-                            <input type="radio" id="strategyTensor" name="strategy" value="tensor">
-                            <label for="strategyTensor">Tensor</label>
+                            <input type="radio" id="shardingTensor" name="sharding" value="Tensor">
+                            <label for="shardingTensor">Tensor</label>
                         </div>
+                    </div>
+                </div>
+                
+                <div class="strategy-selector">
+                    <label class="launch-label">Instance Type:</label>
+                    <div class="strategy-options">
                         <div class="strategy-option">
-                            <input type="radio" id="strategyPipelineRdma" name="strategy" value="pipeline_rdma">
-                            <label for="strategyPipelineRdma">Pipeline RDMA</label>
+                            <input type="radio" id="instanceMlxRing" name="instance_meta" value="MlxRing" checked>
+                            <label for="instanceMlxRing">MLX Ring</label>
                         </div>
                         <div class="strategy-option">
-                            <input type="radio" id="strategyTensorRdma" name="strategy" value="tensor_rdma">
-                            <label for="strategyTensorRdma">Tensor RDMA</label>
+                            <input type="radio" id="instanceMlxIbv" name="instance_meta" value="MlxIbv">
+                            <label for="instanceMlxIbv">MLX IBV</label>
                         </div>
                     </div>
                 </div>
@@ -1277,8 +1279,10 @@
                 return;
             }
 
-            const selectedStrategy = document.querySelector('input[name="strategy"]:checked').value;
-            console.log("selectedStrategy", selectedStrategy);
+            const selectedSharding = document.querySelector('input[name="sharding"]:checked').value;
+            const selectedInstanceMeta = document.querySelector('input[name="instance_meta"]:checked').value;
+            console.log("selectedSharding", selectedSharding);
+            console.log("selectedInstanceMeta", selectedInstanceMeta);
 
             try {
                 showLaunchStatus('Launching instance...', 'loading');
@@ -1291,7 +1295,8 @@
                     },
                     body: JSON.stringify({ 
                         model_id: selectedModelId,
-                        strategy: selectedStrategy
+                        sharding: selectedSharding,
+                        instance_meta: selectedInstanceMeta
                     })
                 });
 
@@ -1333,7 +1338,13 @@
         }
 
         // Calculate download status for an instance based on its runners, with detailed per-file info
-        function calculateInstanceDownloadStatus(instance, runners) {
+        function calculateInstanceDownloadStatus(instanceWrapped, runners) {
+            // Unwrap tagged Instance union (MlxRingInstance or MlxIbvInstance)
+            const [_instanceTag, instance] = getTagged(instanceWrapped);
+            if (!instance || typeof instance !== 'object') {
+                return { isDownloading: false, progress: 0, details: [] };
+            }
+            
             if (!instance.shardAssignments?.runnerToShard || !runners) {
                 return { isDownloading: false, progress: 0, details: [] };
             }
@@ -1423,28 +1434,36 @@
         }
 
 
-        // Derive a display status for an instance from its runners.
-        // Priority: FAILED > DOWNLOADING > STARTING > RUNNING > LOADED > INACTIVE
-        function deriveInstanceStatus(instance, runners = {}) {
-            const runnerIds = Object.keys(instance.shardAssignments?.runnerToShard || {});
+        // Helper function to unwrap tagged unions (defined globally for reuse)
+        function getTagged(obj) {
+            if (!obj || typeof obj !== 'object') return [null, null];
+            const keys = Object.keys(obj);
+            if (keys.length === 1 && typeof keys[0] === 'string') {
+                return [keys[0], obj[keys[0]]];
+            }
+            return [null, null];
+        }
 
-            function getTagged(obj) {
-                if (!obj || typeof obj !== 'object') return [null, null];
-                const keys = Object.keys(obj);
-                if (keys.length === 1 && typeof keys[0] === 'string') {
-                    return [keys[0], obj[keys[0]]];
-                }
-                return [null, null];
+        // Derive a display status for an instance from its runners.
+        // Priority: FAILED > DOWNLOADING > STARTING > RUNNING > READY > LOADED > INACTIVE
+        function deriveInstanceStatus(instanceWrapped, runners = {}) {
+            // Unwrap tagged Instance union
+            const [_instanceTag, instance] = getTagged(instanceWrapped);
+            if (!instance || typeof instance !== 'object') {
+                return { statusText: 'UNKNOWN', statusClass: 'inactive' };
             }
+            
+            const runnerIds = Object.keys(instance.shardAssignments?.runnerToShard || {});
 
             function canonicalStatusFromKind(kind) {
                 const map = {
-                    DownloadingRunnerStatus: 'Downloading',
-                    InactiveRunnerStatus: 'Inactive',
-                    StartingRunnerStatus: 'Starting',
-                    LoadedRunnerStatus: 'Loaded',
-                    RunningRunnerStatus: 'Running',
-                    FailedRunnerStatus: 'Failed',
+                    RunnerWaitingForModel: 'WaitingForModel',
+                    RunnerLoading: 'Loading',
+                    RunnerLoaded: 'Loaded',
+                    RunnerWarmingUp: 'WarmingUp',
+                    RunnerReady: 'Ready',
+                    RunnerRunning: 'Running',
+                    RunnerFailed: 'Failed',
                 };
                 return map[kind] || null;
             }
@@ -1464,32 +1483,24 @@
             const every = (pred) => statuses.length > 0 && statuses.every(pred);
 
             if (statuses.length === 0) {
-                const it = instance.instanceType;
-                const inactive = (it === 'Inactive' || it === 'INACTIVE');
-                return { statusText: inactive ? 'INACTIVE' : 'LOADED', statusClass: inactive ? 'inactive' : 'loaded' };
+                return { statusText: 'UNKNOWN', statusClass: 'inactive' };
             }
 
             if (has('Failed')) return { statusText: 'FAILED', statusClass: 'failed' };
-            if (has('Downloading')) return { statusText: 'DOWNLOADING', statusClass: 'downloading' };
-            if (has('Starting')) return { statusText: 'LOADING', statusClass: 'starting' };
+            if (has('Loading')) return { statusText: 'LOADING', statusClass: 'starting' };
+            if (has('WarmingUp')) return { statusText: 'WARMING UP', statusClass: 'starting' };
             if (has('Running')) return { statusText: 'RUNNING', statusClass: 'running' };
+            if (has('Ready')) return { statusText: 'READY', statusClass: 'loaded' };
+            if (has('Loaded')) return { statusText: 'LOADED', statusClass: 'loaded' };
+            if (has('WaitingForModel')) return { statusText: 'WAITING', statusClass: 'starting' };
 
-            const allInactive = every(s => s === 'Inactive');
-            const loadedOrInactiveOnly = every(s => s === 'Loaded' || s === 'Inactive');
-            const anyLoaded = statuses.some(s => s === 'Loaded');
-            if (loadedOrInactiveOnly && anyLoaded) {
-                return { statusText: 'LOADED', statusClass: 'loaded' };
-            }
-            if (allInactive) {
-                return { statusText: 'INACTIVE', statusClass: 'inactive' };
-            }
-            return { statusText: 'LOADED', statusClass: 'loaded' };
+            return { statusText: 'UNKNOWN', statusClass: 'inactive' };
         }
 
         function renderInstances(instances, runners = {}) {
-            const instancesArray = Object.values(instances);
+            const instanceEntries = Object.entries(instances || {});
             
-            if (instancesArray.length === 0) {
+            if (instanceEntries.length === 0) {
                 instancesList.innerHTML = '<div class="no-instances">No instances running</div>';
                 return;
             }
@@ -1498,8 +1509,17 @@
             instanceIdToColor = {};
             connectionToInstances = {};
             
-            instancesArray.forEach(instance => {
-                const instanceId = instance.instanceId;
+            instanceEntries.forEach(([instanceId, instanceWrapped]) => {
+                // Validate instanceId
+                if (!instanceId || typeof instanceId !== 'string') {
+                    return;
+                }
+                
+                // Unwrap tagged Instance union
+                const [_instanceTag, instance] = getTagged(instanceWrapped);
+                if (!instance || typeof instance !== 'object') {
+                    return;
+                }
                 instanceIdToColor[instanceId] = generateInstanceColor(instanceId);
                 
                 // Determine which nodes this instance uses
@@ -1521,11 +1541,22 @@
                 }
             });
 
-            const instancesHTML = instancesArray.map(instance => {
+            const instancesHTML = instanceEntries.map(([instanceId, instanceWrapped]) => {
+                // Validate instanceId
+                if (!instanceId || typeof instanceId !== 'string') {
+                    return '';
+                }
+                
+                // Unwrap tagged Instance union
+                const [_instanceTag, instance] = getTagged(instanceWrapped);
+                if (!instance || typeof instance !== 'object') {
+                    return '';
+                }
+                
                 const modelId = instance.shardAssignments?.modelId || 'Unknown Model';
-                const truncatedInstanceId = instance.instanceId.length > 8 
-                    ? instance.instanceId.substring(0, 8) + '...' 
-                    : instance.instanceId;
+                const truncatedInstanceId = instanceId.length > 8 
+                    ? instanceId.substring(0, 8) + '...' 
+                    : instanceId;
                 
                 // Create reverse mapping from runnerId to nodeId using nodeToRunner
                 const nodeToRunner = instance.shardAssignments?.nodeToRunner || {};
@@ -1534,20 +1565,31 @@
                     runnerToNode[runnerId] = nodeId;
                 });
                 
-                // Extract parallelization strategy from the first shard
+                // Extract sharding strategy from the first shard
+                // Shards are tagged unions: {"PipelineShardMetadata": {...}} or {"TensorShardMetadata": {...}}
                 const runnerToShard = instance.shardAssignments?.runnerToShard || {};
-                const firstShardData = Object.values(runnerToShard)[0];
-                let parallelizationStrategy = 'Unknown';
-                if (firstShardData) {
-                    const shardKeys = Object.keys(firstShardData);
-                    if (shardKeys.length === 1) {
-                        const shardPayload = firstShardData[shardKeys[0]];
-                        parallelizationStrategy = shardPayload?.strategy || firstShardData.strategy || 'Unknown';
-                    } else {
-                        parallelizationStrategy = firstShardData.strategy || 'Unknown';
+                const firstShardWrapped = Object.values(runnerToShard)[0];
+                let shardingType = 'Unknown';
+                if (firstShardWrapped) {
+                    const [shardTag, _shardData] = getTagged(firstShardWrapped);
+                    if (shardTag === 'PipelineShardMetadata') {
+                        shardingType = 'Pipeline';
+                    } else if (shardTag === 'TensorShardMetadata') {
+                        shardingType = 'Tensor';
                     }
                 }
                 
+                // Extract instance type from the tagged union
+                // Instance is tagged as {"MlxRingInstance": {...}} or {"MlxIbvInstance": {...}}
+                let instanceType = 'Unknown';
+                if (_instanceTag === 'MlxRingInstance') {
+                    instanceType = 'MLX Ring';
+                } else if (_instanceTag === 'MlxIbvInstance') {
+                    instanceType = 'MLX IBV';
+                }
+                
+                const parallelizationStrategy = `${shardingType} (${instanceType})`;
+                
                 // Generate hosts HTML using runner IDs and friendly names
                 const runnerIds = Object.keys(runnerToShard);
                 const hostsHTML = runnerIds.map(runnerId => {
@@ -1559,14 +1601,14 @@
                     return `<span class="instance-host">${friendlyName} (${shortId})</span>`;
                 }).join('') || '';
 
-                // Calculate download status for this instance
-                const downloadStatus = calculateInstanceDownloadStatus(instance, runners);
+                // Calculate download status for this instance (pass wrapped instance)
+                const downloadStatus = calculateInstanceDownloadStatus(instanceWrapped, runners);
 
                 let statusText, statusClass;
                 if (downloadStatus.isDownloading) {
                     ({ statusText, statusClass } = { statusText: 'DOWNLOADING', statusClass: 'downloading' });
                 } else {
-                    ({ statusText, statusClass } = deriveInstanceStatus(instance, runners));
+                    ({ statusText, statusClass } = deriveInstanceStatus(instanceWrapped, runners));
                 }
 
                 // Generate download progress HTML - overall + per node with file details
@@ -1660,7 +1702,7 @@
                 const shardCount = Object.keys(runnerToShard).length;
                 
                 // Use the instance's color for the indicator
-                const instanceColor = instanceIdToColor[instance.instanceId] || 'var(--exo-yellow)';
+                const instanceColor = instanceIdToColor[instanceId] || 'var(--exo-yellow)';
                 const borderStyle = `background-color: ${instanceColor};`;
                 
                 return `
@@ -1672,7 +1714,7 @@
                                 <span class="instance-status ${statusClass}">${statusText}</span>
                             </div>
                             <div class="instance-actions">
-                                <button class="instance-delete-button" data-instance-id="${instance.instanceId}" title="Delete Instance">
+                                <button class="instance-delete-button" data-instance-id="${instanceId}" title="Delete Instance">
                                     Delete
                                 </button>
                             </div>
diff --git a/flake.nix b/flake.nix
index 2d70276a..d2bd1b67 100644
--- a/flake.nix
+++ b/flake.nix
@@ -58,8 +58,9 @@
               # NIX
               nixpkgs-fmt
 
-              # JUST
+              # MISC
               just
+              jq
             ]
             ++ (pkgs.lib.optionals pkgs.stdenv.isLinux [
               # IFCONFIG
diff --git a/pyproject.toml b/pyproject.toml
index 12ff2bdf..5a7f8fa9 100644
--- a/pyproject.toml
+++ b/pyproject.toml
@@ -32,7 +32,7 @@ dependencies = [
     "loguru>=0.7.3",
     "textual>=5.3.0",
     "exo_pyo3_bindings", # rust bindings
-    "anyio>=4.10.0",
+    "anyio>=4.11.0",
     "bidict>=0.23.1",
     "mlx>=0.29.3",
     "mlx-lm>=0.28.3",
diff --git a/src/exo/engines/mlx/auto_parallel.py b/src/exo/engines/mlx/auto_parallel.py
index d1026779..3223f86f 100644
--- a/src/exo/engines/mlx/auto_parallel.py
+++ b/src/exo/engines/mlx/auto_parallel.py
@@ -1,7 +1,9 @@
 from abc import ABC, abstractmethod
 from functools import partial
+from inspect import signature
 from typing import TYPE_CHECKING, Callable, Protocol, cast, override
 
+from mlx_lm.models.cache import KVCache
 from mlx_lm.models.deepseek_v3 import DeepseekV3MLP
 from mlx_lm.models.deepseek_v3 import Model as DeepseekV3Model
 from mlx_lm.models.llama import Model as LlamaModel
@@ -12,8 +14,6 @@ import mlx.core as mx
 import mlx.nn as nn
 from exo.shared.types.worker.shards import (
     PipelineShardMetadata,
-    ShardMetadata,
-    TensorShardMetadata,
 )
 from mlx.nn.layers.distributed import (
     shard_inplace,
@@ -58,12 +58,10 @@ class PipelineFirstLayer(CustomMlxLayer):
         self,
         original_layer: _LayerCallable,
         r: int,
-        s: int,
-        group: mx.distributed.Group | None = None,
+        group: mx.distributed.Group,
     ):
         super().__init__(original_layer)
         self.r: int = r
-        self.s: int = s
         self.group = group
 
     @override
@@ -79,177 +77,167 @@ class PipelineLastLayer(CustomMlxLayer):
         original_layer: _LayerCallable,
         r: int,
         s: int,
-        group: mx.distributed.Group | None = None,
+        group: mx.distributed.Group,
     ):
         super().__init__(original_layer)
         self.r: int = r
         self.s: int = s
         self.group = group
+        self.original_layer_signature = signature(self.original_layer.__call__)
 
     @override
     def __call__(
-        self, x: mx.array, *args: object, cache: object = None, **kwargs: object
+        self, x: mx.array, *args: object, **kwargs: object
     ) -> mx.array:
+
+        cache = self.original_layer_signature.bind_partial(x, *args, **kwargs).arguments.get("cache", None)
+
+        assert cache is None or isinstance(cache, KVCache)
+
         output: mx.array = self.original_layer(x, *args, **kwargs)
+
         if self.r != self.s - 1:
             output = mx.distributed.send(
                 output, (self.r + 1) % self.s, group=self.group
             )
             if (
-                cache is not None
-                and hasattr(cache, "keys")
-                and getattr(cache, "keys", None) is not None
-            ):
-                cache.keys = mx.depends(cache.keys, output)  # type: ignore[reportUnknownMemberType]
+                    cache is not None
+                    and hasattr(cache, "keys")
+                    and getattr(cache, "keys", None) is not None
+                ):
+                    # This change happened upstream - check out mlx github somewhere??
+                    cache.keys = mx.depends(cache.keys, output)  # type: ignore[reportUnknownMemberType]
 
         output = mx.distributed.all_gather(output, group=self.group)[-output.shape[0] :]
         return output
 
 
-class ParallelisationShardStrategy(Protocol):
-    def auto_parallel(
-        self, model: nn.Module, model_shard_meta: ShardMetadata
-    ) -> nn.Module: ...
-
-
-class PipelineParallelisationStrategy(ParallelisationShardStrategy):
-    def __init__(self, group: mx.distributed.Group):
-        self.group = group
-
-    def auto_parallel(
-        self, model: nn.Module, model_shard_meta: ShardMetadata
-    ) -> nn.Module:
-        """
-        Automatically parallelize a model across multiple devices.
-        Args:
-        model: The model to parallelize (must have a 'layers' or 'h' property)
-        model_shard_meta: The metadata for the model shard
-        Returns:
-        The parallelized model
-        """
-        assert isinstance(model_shard_meta, PipelineShardMetadata)
-
-        inner_model_instance: nn.Module = PipelineParallelisationStrategy._inner_model(
-            model
-        )
-
-        # Handle both model.layers and model.h cases
-        layers: list[_LayerCallable]
-        if hasattr(inner_model_instance, "layers"):
-            layers = cast(list[_LayerCallable], inner_model_instance.layers)
-        elif hasattr(inner_model_instance, "h"):
-            layers = cast(list[_LayerCallable], inner_model_instance.h)
-        else:
-            raise ValueError("Model must have either a 'layers' or 'h' attribute")
-
-        layers = layers[model_shard_meta.start_layer : model_shard_meta.end_layer]
-        layers[0] = PipelineFirstLayer(
-            layers[0],
-            model_shard_meta.device_rank,
-            model_shard_meta.world_size,
-            group=self.group,
-        )
-        layers[-1] = PipelineLastLayer(
-            layers[-1],
-            model_shard_meta.device_rank,
-            model_shard_meta.world_size,
-            group=self.group,
-        )
-
-        PipelineParallelisationStrategy._set_layers(model, layers)
-
-        return model
-
-    @staticmethod
-    def _inner_model(model: nn.Module) -> nn.Module:
-        inner = getattr(model, "model", None)
-        if isinstance(inner, nn.Module):
-            return inner
-
-        inner = getattr(model, "transformer", None)
-        if isinstance(inner, nn.Module):
-            return inner
-
-        raise ValueError("Model must either have a 'model' or 'transformer' attribute")
-
-    @staticmethod
-    def _set_layers(model: nn.Module, layers: list[_LayerCallable]) -> None:
-        inner_model_instance = PipelineParallelisationStrategy._inner_model(model)
-        if hasattr(inner_model_instance, "layers"):
-            inner_model_instance.layers = layers
-
-            # Update DeepSeek V3 specific parameters when layers are shrunk
-            if isinstance(model, DeepseekV3Model) and hasattr(
-                inner_model_instance, "num_layers"
-            ):
-                inner_model_instance.start_idx = 0
-                inner_model_instance.end_idx = len(layers)
-                inner_model_instance.num_layers = len(layers)
-        elif hasattr(inner_model_instance, "h"):
-            inner_model_instance.h = layers
-        else:
-            raise ValueError("Model must have either a 'layers' or 'h' attribute")
-
-
-class TensorParallelisationStrategy(ParallelisationShardStrategy):
-    def __init__(self, group: mx.distributed.Group):
-        self.group = group
-
-    def auto_parallel(
-        self, model: nn.Module, model_shard_meta: ShardMetadata
-    ) -> nn.Module:
-        assert isinstance(model_shard_meta, TensorShardMetadata)
-
-        all_to_sharded_linear = partial(
-            shard_linear,
-            sharding="all-to-sharded",
-            group=self.group,
-        )
-        sharded_to_all_linear = partial(
-            shard_linear,
-            sharding="sharded-to-all",
-            group=self.group,
+def _set_layers(model: nn.Module, layers: list[_LayerCallable]) -> None:
+    inner_model_instance = _inner_model(model)
+    if hasattr(inner_model_instance, "layers"):
+        inner_model_instance.layers = layers
+
+        # Update DeepSeek V3 specific parameters when layers are shrunk
+        if isinstance(model, DeepseekV3Model) and hasattr(
+            inner_model_instance, "num_layers"
+        ):
+            inner_model_instance.start_idx = 0
+            inner_model_instance.end_idx = len(layers)
+            inner_model_instance.num_layers = len(layers)
+    elif hasattr(inner_model_instance, "h"):
+        inner_model_instance.h = layers
+    else:
+        raise ValueError("Model must have either a 'layers' or 'h' attribute")
+
+
+def pipeline_auto_parallel(
+    model: nn.Module,
+    group: mx.distributed.Group,
+    model_shard_meta: PipelineShardMetadata,
+) -> nn.Module:
+    """
+    Automatically parallelize a model across multiple devices.
+    Args:
+    model: The model to parallelize (must have a 'layers' or 'h' property)
+    model_shard_meta: The metadata for the model shard
+    Returns:
+    The parallelized model
+    """
+    inner_model_instance: nn.Module = _inner_model(model)
+
+    # Handle both model.layers and model.h cases
+    layers: list[_LayerCallable]
+    if hasattr(inner_model_instance, "layers"):
+        layers = cast(list[_LayerCallable], inner_model_instance.layers)
+    elif hasattr(inner_model_instance, "h"):
+        layers = cast(list[_LayerCallable], inner_model_instance.h)
+    else:
+        raise ValueError("Model must have either a 'layers' or 'h' attribute")
+
+    layers = layers[model_shard_meta.start_layer : model_shard_meta.end_layer]
+    layers[0] = PipelineFirstLayer(layers[0], model_shard_meta.device_rank, group=group)
+    layers[-1] = PipelineLastLayer(
+        layers[-1],
+        model_shard_meta.device_rank,
+        model_shard_meta.world_size,
+        group=group,
+    )
+
+    _set_layers(model, layers)
+
+    assert isinstance(layers, list), (
+        "Expected a list of layers after auto-parallel initialisation"
+    )
+
+    return model
+
+
+def _inner_model(model: nn.Module) -> nn.Module:
+    inner = getattr(model, "model", None)
+    if isinstance(inner, nn.Module):
+        return inner
+
+    inner = getattr(model, "transformer", None)
+    if isinstance(inner, nn.Module):
+        return inner
+
+    raise ValueError("Model must either have a 'model' or 'transformer' attribute")
+
+
+def tensor_auto_parallel(
+    model: nn.Module,
+    group: mx.distributed.Group,
+) -> nn.Module:
+    all_to_sharded_linear = partial(
+        shard_linear,
+        sharding="all-to-sharded",
+        group=group,
+    )
+    sharded_to_all_linear = partial(
+        shard_linear,
+        sharding="sharded-to-all",
+        group=group,
+    )
+
+    all_to_sharded_linear_in_place = partial(
+        shard_inplace,
+        sharding="all-to-sharded",
+        group=group,
+    )
+    sharded_to_all_linear_in_place = partial(
+        shard_inplace,
+        sharding="sharded-to-all",
+        group=group,
+    )
+
+    if isinstance(model, LlamaModel):
+        tensor_parallel_sharding_strategy = LlamaShardingStrategy(
+            group,
+            all_to_sharded_linear,
+            sharded_to_all_linear,
+            all_to_sharded_linear_in_place,
+            sharded_to_all_linear_in_place,
         )
-
-        all_to_sharded_linear_in_place = partial(
-            shard_inplace,
-            sharding="all-to-sharded",
-            group=self.group,
+    elif isinstance(model, DeepseekV3Model):
+        tensor_parallel_sharding_strategy = DeepSeekShardingStrategy(
+            group,
+            all_to_sharded_linear,
+            sharded_to_all_linear,
+            all_to_sharded_linear_in_place,
+            sharded_to_all_linear_in_place,
         )
-        sharded_to_all_linear_in_place = partial(
-            shard_inplace,
-            sharding="sharded-to-all",
-            group=self.group,
+    elif isinstance(model, Qwen3MoeModel):
+        tensor_parallel_sharding_strategy = QwenShardingStrategy(
+            group,
+            all_to_sharded_linear,
+            sharded_to_all_linear,
+            all_to_sharded_linear_in_place,
+            sharded_to_all_linear_in_place,
         )
+    else:
+        raise ValueError(f"Unsupported model type: {type(model)}")
 
-        if isinstance(model, LlamaModel):
-            tensor_parallel_sharding_strategy = LlamaShardingStrategy(
-                self.group,
-                all_to_sharded_linear,
-                sharded_to_all_linear,
-                all_to_sharded_linear_in_place,
-                sharded_to_all_linear_in_place,
-            )
-        elif isinstance(model, DeepseekV3Model):
-            tensor_parallel_sharding_strategy = DeepSeekShardingStrategy(
-                self.group,
-                all_to_sharded_linear,
-                sharded_to_all_linear,
-                all_to_sharded_linear_in_place,
-                sharded_to_all_linear_in_place,
-            )
-        elif isinstance(model, Qwen3MoeModel):
-            tensor_parallel_sharding_strategy = QwenShardingStrategy(
-                self.group,
-                all_to_sharded_linear,
-                sharded_to_all_linear,
-                all_to_sharded_linear_in_place,
-                sharded_to_all_linear_in_place,
-            )
-        else:
-            raise ValueError(f"Unsupported model type: {type(model)}")
-
-        return tensor_parallel_sharding_strategy.shard_model(model)
+    return tensor_parallel_sharding_strategy.shard_model(model)
 
 
 class TensorParallelShardingStrategy(ABC):
diff --git a/src/exo/engines/mlx/utils_mlx.py b/src/exo/engines/mlx/utils_mlx.py
index 9dfc5df5..e8e6391b 100644
--- a/src/exo/engines/mlx/utils_mlx.py
+++ b/src/exo/engines/mlx/utils_mlx.py
@@ -1,12 +1,10 @@
-import asyncio
-import concurrent.futures
 import os
 import resource
-from asyncio import AbstractEventLoop
 from typing import Any, Callable, cast
 
 from mlx_lm.models.cache import KVCache
 from mlx_lm.sample_utils import make_sampler
+from mlx_lm.tokenizer_utils import TokenizerWrapper
 
 from exo.worker.runner.utils import get_weights_size
 
@@ -15,22 +13,30 @@ try:
 except ImportError:
     from mlx_lm.tokenizer_utils import load as load_tokenizer  # type: ignore
 from mlx_lm.utils import load_model
-from mlx.utils import tree_reduce
 from pydantic import RootModel
 
 import mlx.core as mx
 import mlx.nn as nn
-from exo.engines.mlx import Model, TokenizerWrapper
+from exo.engines.mlx import Model
 from exo.engines.mlx.auto_parallel import (
-    PipelineParallelisationStrategy,
-    TensorParallelisationStrategy,
+    pipeline_auto_parallel,
+    tensor_auto_parallel,
 )
 from exo.shared.types.memory import Memory
 from exo.shared.types.common import Host
 from exo.shared.types.tasks import ChatCompletionTaskParams
-from exo.shared.types.worker.communication import runner_print
-from exo.shared.types.worker.shards import ShardMetadata
+from exo.shared.types.worker.instances import (
+    BoundInstance,
+    MlxIbvInstance,
+    MlxRingInstance,
+)
+from exo.shared.types.worker.shards import (
+    PipelineShardMetadata,
+    ShardMetadata,
+    TensorShardMetadata,
+)
 from exo.worker.download.download_utils import build_model_path
+from exo.worker.runner.bootstrap import logger
 
 # Needed for 8 bit model
 resource.setrlimit(resource.RLIMIT_NOFILE, (2048, 4096))
@@ -70,11 +76,7 @@ class HostList(RootModel[list[str]]):
 
 
 def mlx_distributed_init(
-    rank: int,
-    hosts: list[Host] | None = None,
-    mlx_ibv_devices: list[list[str | None]] | None = None,
-    mlx_ibv_coordinator: str | None = None,
-    strict: bool = True,
+    bound_instance: BoundInstance,
 ) -> mx.distributed.Group:
     """
     Initialize the MLX distributed (runs in thread pool).
@@ -85,117 +87,100 @@ def mlx_distributed_init(
     - mlx_ibv_coordinator: coordinator address (IP:PORT) for RDMA setup
     - strict: if True, raise an error if the distributed backend is not available
     """
-    runner_print(f"Starting initialization for rank {rank}. Strict: {strict}")
+    rank = bound_instance.bound_shard().device_rank
+    logger.info(f"Starting initialization for rank {rank}")
 
-    if mlx_ibv_devices is not None and mlx_ibv_devices != []:
-        assert mlx_ibv_coordinator is not None, (
-            "To use ibv backend must set ibv coordinator"
-        )
-        import json
+    # TODO: singleton instances
+    match bound_instance.instance:
+        case MlxRingInstance(hosts=hosts):
+            hostfile = f"./hosts_{rank}.json"
+            hosts_json = HostList.from_hosts(hosts).model_dump_json()
 
-        # Use RDMA connectivity matrix
-        devices_file = f"./hosts_{rank}.json"
-        ibv_devices_json = json.dumps(mlx_ibv_devices)
-        runner_print(f"rank {rank} MLX_IBV_DEVICES: {ibv_devices_json}")
-        runner_print(f"rank {rank} MLX_IBV_COORDINATOR: {mlx_ibv_coordinator}")
+            with open(hostfile, "w") as f:
+                _ = f.write(hosts_json)
 
-        with open(devices_file, "w") as f:
-            _ = f.write(ibv_devices_json)
+            logger.info(f"rank {rank} hostfile: {hostfile} hosts: {hosts_json}")
 
-        os.environ["MLX_IBV_DEVICES"] = devices_file
-        os.environ["MLX_RANK"] = str(rank)
-        os.environ["MLX_IBV_COORDINATOR"] = mlx_ibv_coordinator
-    elif hosts is not None and hosts != []:
-        # Traditional host-based connectivity
-        hostfile = f"./hosts_{rank}.json"
-        hosts_json = HostList.from_hosts(hosts).model_dump_json()
+            os.environ["MLX_HOSTFILE"] = hostfile
+            os.environ["MLX_RANK"] = str(rank)
+            os.environ["MLX_RING_VERBOSE"] = "1"
+            group = mx.distributed.init(backend="ring", strict=True)
 
-        runner_print(f"rank {rank} hostfile: {hostfile} hosts: {hosts_json}")
+        case MlxIbvInstance(ibv_devices=ibv_devices, ibv_coordinator=ibv_coordinator):
+            import json
 
-        with open(hostfile, "w") as f:
-            _ = f.write(hosts_json)
+            # Use RDMA connectivity matrix
+            devices_file = f"./hosts_{rank}.json"
+            ibv_devices_json = json.dumps(ibv_devices)
 
-        os.environ["MLX_HOSTFILE"] = hostfile
-        os.environ["MLX_RANK"] = str(rank)
-        os.environ["MLX_RING_VERBOSE"] = "1"
-    else:
-        runner_print("No distributed setup, using single device mode")
+            with open(devices_file, "w") as f:
+                _ = f.write(ibv_devices_json)
 
-    group = mx.distributed.init(
-        backend="ring" if hosts is not None else "ibv",
-        strict=strict,
-    )
-    runner_print(f"Rank {rank} mlx distributed initialization complete")
+            logger.info(f"rank {rank} MLX_IBV_DEVICES: {ibv_devices_json}")
+            logger.info(f"rank {rank} MLX_IBV_COORDINATOR: {ibv_coordinator}")
+            os.environ["MLX_IBV_DEVICES"] = devices_file
+            os.environ["MLX_RANK"] = str(rank)
+            os.environ["MLX_IBV_COORDINATOR"] = ibv_coordinator
+            group = mx.distributed.init(backend="ibv", strict=True)
+
+    logger.info(f"Rank {rank} mlx distributed initialization complete")
 
     return group
 
 
 def initialize_mlx(
-    model_shard_meta: ShardMetadata,
-    hosts: list[Host] | None = None,
-    mlx_ibv_devices: list[list[str | None]] | None = None,
-    mlx_ibv_coordinator: str | None = None,
-) -> tuple[
-    Model, TokenizerWrapper, Callable[[mx.array], mx.array], mx.distributed.Group
-]:
+    bound_instance: BoundInstance,
+) -> tuple[Model, TokenizerWrapper, Callable[[mx.array], mx.array]]:
     """
     Initialize the MLX model, tokenizer, and sampler. Runs in the MLX thread.
-
-    Either hosts or mlx_ibv_devices must be provided for distributed setups:
-    - hosts: traditional host-based connectivity
-    - mlx_ibv_devices: RDMA connectivity matrix
     """
     mx.random.seed(42)
-    group = mlx_distributed_init(
-        model_shard_meta.device_rank,
-        hosts=hosts,
-        mlx_ibv_devices=mlx_ibv_devices,
-        mlx_ibv_coordinator=mlx_ibv_coordinator,
-        strict=(hosts is not None and len(hosts) > 1)
-        or (mlx_ibv_devices is not None and len(mlx_ibv_devices) > 1),
-    )
 
-    set_wired_limit_for_model(get_weights_size(model_shard_meta))
+    set_wired_limit_for_model(get_weights_size(bound_instance.bound_shard()))
 
     sampler: Callable[[mx.array], mx.array] = make_sampler(temp=0.7)
+    logger.info("Created a sampler")
 
-    model, tokenizer = shard_and_load(model_shard_meta, group=group)
-    model = cast(Model, model)
+    if len(bound_instance.instance.shard_assignments.node_to_runner) <= 1:
+        logger.info(f"Single device used for {bound_instance.instance}")
+        model_path = build_model_path(bound_instance.bound_shard().model_meta.model_id)
+        model, _ = load_model(model_path, strict=True)
+        # TODO: we should really make this opt-in, but Kimi requires trust_remote_code=True
+        tokenizer = cast(TokenizerWrapper, load_tokenizer(model_path, tokenizer_config_extra={"trust_remote_code": True}))
+        assert isinstance(tokenizer, TokenizerWrapper)
 
-    return model, tokenizer, sampler, group
+    else:
+        logger.info("Starting distributed init")
+        group = mlx_distributed_init(bound_instance)
+        model, tokenizer = shard_and_load(bound_instance.bound_shard(), group=group)
+
+    set_wired_limit_for_model(get_weights_size(bound_instance.bound_shard()))
+
+    return cast(Model, model), tokenizer, sampler
 
 
 def shard_and_load(
-    model_shard_meta: ShardMetadata,
+    shard_metadata: ShardMetadata,
     group: mx.distributed.Group,
 ) -> tuple[nn.Module, TokenizerWrapper]:
-    model_path = build_model_path(model_shard_meta.model_meta.model_id)
-
-    runner_print(
-        f"loading model from {model_path} with strategy {model_shard_meta.strategy}"
-    )
+    model_path = build_model_path(shard_metadata.model_meta.model_id)
 
     model, config = load_model(model_path, lazy=True, strict=False)
-    runner_print(f"{config=}")
+    logger.info(f"{config=}")
     assert isinstance(model, nn.Module)
 
-    tokenizer = cast(TokenizerWrapper, load_tokenizer(model_path))
-
-    runner_print(f"Group size: {group.size()}, group rank: {group.rank()}")
+    # TODO: we should really make this opt-in, but Kimi requires trust_remote_code=True
+    tokenizer = cast(TokenizerWrapper, load_tokenizer(model_path, tokenizer_config_extra={"trust_remote_code": True}))
 
-    match model_shard_meta.strategy:
-        case "auto":
-            strategy = PipelineParallelisationStrategy(group)
-        case "pipeline":
-            strategy = PipelineParallelisationStrategy(group)
-        case "pipeline_rdma":
-            strategy = PipelineParallelisationStrategy(group)
-        case "tensor":
-            strategy = TensorParallelisationStrategy(group)
-        case "tensor_rdma":
-            strategy = TensorParallelisationStrategy(group)
+    logger.info(f"Group size: {group.size()}, group rank: {group.rank()}")
 
-    model = strategy.auto_parallel(model, model_shard_meta)
+    match shard_metadata:
+        case TensorShardMetadata():
+            logger.info(f"loading model from {model_path} with tensor parallelism")
+            model = tensor_auto_parallel(model, group)
+        case PipelineShardMetadata():
+            logger.info(f"loading model from {model_path} with pipeline parallelism")
+            model = pipeline_auto_parallel(model, group, shard_metadata)
 
     mx.eval(model.parameters())
     mx.eval(model)
@@ -206,13 +191,10 @@ def shard_and_load(
     return model, tokenizer
 
 
-async def apply_chat_template(
-    mlx_executor: concurrent.futures.ThreadPoolExecutor,
+def apply_chat_template(
     tokenizer: TokenizerWrapper,
     chat_task_data: ChatCompletionTaskParams,
 ) -> str:
-    loop: AbstractEventLoop = asyncio.get_running_loop()
-
     # Now we can properly access the messages
     messages = chat_task_data.messages
     messages_dicts: list[dict[str, Any]] = [msg.model_dump() for msg in messages]
@@ -237,16 +219,13 @@ async def apply_chat_template(
 
     messages_dicts = formatted_messages
 
-    prompt: str = await loop.run_in_executor(
-        executor=mlx_executor,
-        func=lambda: tokenizer.apply_chat_template(
-            messages_dicts,
-            tokenize=False,
-            add_generation_prompt=True,
-        ),
+    prompt: str = tokenizer.apply_chat_template(  # type: ignore
+        messages_dicts,
+        tokenize=False,
+        add_generation_prompt=True,
     )
 
-    return prompt
+    return prompt  # type: ignore
 
 
 class NullKVCache(KVCache):
@@ -272,7 +251,7 @@ class NullKVCache(KVCache):
         raise NotImplementedError("We should not be setting a NullKVCache.")
 
 
-async def make_kv_cache(
+def make_kv_cache(
     model: Model,
     max_kv_size: int | None = None,
 ) -> list[KVCache]:
@@ -311,8 +290,8 @@ def set_wired_limit_for_model(model_size: Memory):
     if model_bytes > 0.9 * max_rec_size:
         model_mb = model_bytes // 2**20
         max_rec_mb = max_rec_size // 2**20
-        runner_print(
-            f"[WARNING] Generating with a model that requires {model_mb} MB "
+        logger.warning(
+            f"Generating with a model that requires {model_mb} MB "
             f"which is close to the maximum recommended size of {max_rec_mb} "
             "MB. This can be slow. See the documentation for possible work-arounds: "
             "https://github.com/ml-explore/mlx-lm/tree/main#large-models"
@@ -322,4 +301,6 @@ def set_wired_limit_for_model(model_size: Memory):
     target_cache = min(target_cache, max_rec_size)
     mx.set_cache_limit(target_cache)
     mx.set_wired_limit(max_rec_size)
-    runner_print(f"Wired limit set to {max_rec_size}. Cache limit set to {target_cache}.")
+    logger.info(
+        f"Wired limit set to {max_rec_size}. Cache limit set to {target_cache}."
+    )
diff --git a/src/exo/main.py b/src/exo/main.py
index f1496f08..28530879 100644
--- a/src/exo/main.py
+++ b/src/exo/main.py
@@ -1,4 +1,5 @@
 import argparse
+import multiprocessing as mp
 from dataclasses import dataclass
 from typing import Self
 
@@ -15,6 +16,7 @@ from exo.shared.constants import EXO_LOG
 from exo.shared.election import Election, ElectionResult
 from exo.shared.logging import logger_cleanup, logger_setup
 from exo.shared.types.common import NodeId, SessionId
+from exo.shared.types.commands import KillCommand
 from exo.utils.channels import Receiver, channel
 from exo.utils.pydantic_ext import CamelCaseModel
 from exo.worker.download.impl_shard_downloader import exo_shard_downloader
@@ -108,6 +110,17 @@ class Node:
             if self.api:
                 tg.start_soon(self.api.run)
             tg.start_soon(self._elect_loop)
+            tg.start_soon(self._listen_for_kill_command)
+
+    async def _listen_for_kill_command(self):
+        assert self._tg
+        with self.router.receiver(topics.COMMANDS) as commands:
+            async for command in commands:
+                match command.command:
+                    case KillCommand():
+                        self._tg.cancel_scope.cancel()
+                    case _:
+                        pass
 
     async def _elect_loop(self):
         assert self._tg
@@ -185,6 +198,7 @@ class Node:
 
 def main():
     args = Args.parse()
+    mp.set_start_method("spawn")
     # TODO: Refactor the current verbosity system
     logger_setup(EXO_LOG, args.verbosity)
     logger.info("Starting EXO")
diff --git a/src/exo/master/api.py b/src/exo/master/api.py
index a3a3e1fb..b52dbe29 100644
--- a/src/exo/master/api.py
+++ b/src/exo/master/api.py
@@ -35,23 +35,25 @@ from exo.shared.types.commands import (
     CreateInstance,
     DeleteInstance,
     ForwarderCommand,
-    # TODO: SpinUpInstance
+    KillCommand,
     TaskFinished,
 )
 from exo.shared.types.common import CommandId, NodeId, SessionId
 from exo.shared.types.events import ChunkGenerated, Event, ForwarderEvent, IndexedEvent
+from exo.shared.types.memory import Memory
 from exo.shared.types.models import ModelMetadata
 from exo.shared.types.state import State
 from exo.shared.types.tasks import ChatCompletionTaskParams
-from exo.shared.types.worker.common import InstanceId
-from exo.shared.types.worker.instances import Instance
+from exo.shared.types.worker.instances import Instance, InstanceId
 from exo.utils.channels import Receiver, Sender
 from exo.utils.event_buffer import OrderedBuffer
 
 
-def chunk_to_response(chunk: TokenChunk) -> ChatCompletionResponse:
+def chunk_to_response(
+    chunk: TokenChunk, command_id: CommandId
+) -> ChatCompletionResponse:
     return ChatCompletionResponse(
-        id=chunk.command_id,
+        id=command_id,
         created=int(time.time()),
         model=chunk.model,
         choices=[
@@ -117,9 +119,7 @@ class API:
             name="dashboard",
         )
 
-        self._chat_completion_queues: dict[
-            CommandId, asyncio.Queue[ChunkGenerated]
-        ] = {}
+        self._chat_completion_queues: dict[CommandId, asyncio.Queue[TokenChunk]] = {}
         self._tg: TaskGroup | None = None
 
     def reset(self, new_session_id: SessionId, result_clock: int):
@@ -152,23 +152,29 @@ class API:
         self.app.get("/v1/models")(self.get_models)
         self.app.post("/v1/chat/completions")(self.chat_completions)
         self.app.get("/state")(lambda: self.state)
+        self.app.delete("/kill")(self.kill_exo)
+
+    async def kill_exo(self):
+        await self._send(KillCommand())
 
     async def create_instance(
         self, payload: CreateInstanceTaskParams
     ) -> CreateInstanceResponse:
         model_meta = await resolve_model_meta(payload.model_id)
-        strategy = payload.strategy
-        required_memory_bytes = model_meta.storage_size.in_kb
-        available_memory_bytes = self._calculate_total_available_memory()
+        required_memory = model_meta.storage_size
+        available_memory = self._calculate_total_available_memory()
 
-        if required_memory_bytes > available_memory_bytes:
+        if required_memory > available_memory:
             raise HTTPException(
                 status_code=400,
-                detail=f"Insufficient memory to create instance. Required: {required_memory_bytes // (1024**3):.1f}GB, Available: {available_memory_bytes // (1024**3):.1f}GB",
+                detail=f"Insufficient memory to create instance. Required: {required_memory.in_gb:.1f}GB, Available: {available_memory.in_gb:.1f}GB",
             )
 
         command = CreateInstance(
-            command_id=CommandId(), model_meta=model_meta, strategy=strategy
+            command_id=CommandId(),
+            model_meta=model_meta,
+            instance_meta=payload.instance_meta,
+            sharding=payload.sharding,
         )
         await self._send(command)
 
@@ -211,15 +217,16 @@ class API:
             chunk = await asyncio.wait_for(
                 self._chat_completion_queues[command_id].get(), timeout=60
             )
-            if chunk.command_id == command_id:
-                assert isinstance(chunk.chunk, TokenChunk)
-                chunk_response: ChatCompletionResponse = chunk_to_response(chunk.chunk)
-                logger.debug(f"chunk_response: {chunk_response}")
-                yield f"data: {chunk_response.model_dump_json()}\n\n"
+            assert isinstance(chunk, TokenChunk)
+            chunk_response: ChatCompletionResponse = chunk_to_response(
+                chunk, command_id
+            )
+            logger.debug(f"chunk_response: {chunk_response}")
+            yield f"data: {chunk_response.model_dump_json()}\n\n"
 
-                if chunk.chunk.finish_reason is not None:
-                    yield "data: [DONE]\n\n"
-                    finished = True
+            if chunk.finish_reason is not None:
+                yield "data: [DONE]\n\n"
+                finished = True
 
         command = TaskFinished(finished_command_id=command_id)
         await self._send(command)
@@ -281,13 +288,13 @@ class API:
             media_type="text/event-stream",
         )
 
-    def _calculate_total_available_memory(self) -> int:
+    def _calculate_total_available_memory(self) -> Memory:
         """Calculate total available memory across all nodes in bytes."""
-        total_available = 0
+        total_available = Memory()
 
         for node in self.state.topology.list_nodes():
             if node.node_profile is not None:
-                total_available += node.node_profile.memory.ram_available.in_bytes
+                total_available += node.node_profile.memory.ram_available
 
         return total_available
 
@@ -323,15 +330,18 @@ class API:
 
     async def _apply_state(self):
         with self.global_event_receiver as events:
-            async for event in events:
-                self.event_buffer.ingest(event.origin_idx, event.event)
+            async for f_event in events:
+                self.event_buffer.ingest(f_event.origin_idx, f_event.event)
                 for idx, event in self.event_buffer.drain_indexed():
                     self.state = apply(self.state, IndexedEvent(event=event, idx=idx))
                     if (
                         isinstance(event, ChunkGenerated)
                         and event.command_id in self._chat_completion_queues
                     ):
-                        self._chat_completion_queues[event.command_id].put_nowait(event)
+                        assert isinstance(event.chunk, TokenChunk)
+                        self._chat_completion_queues[event.command_id].put_nowait(
+                            event.chunk
+                        )
 
     async def _pause_on_new_election(self):
         with self.election_receiver as ems:
diff --git a/src/exo/master/main.py b/src/exo/master/main.py
index 60285001..7badeeca 100644
--- a/src/exo/master/main.py
+++ b/src/exo/master/main.py
@@ -31,8 +31,14 @@ from exo.shared.types.events import (
     TopologyEdgeDeleted,
 )
 from exo.shared.types.state import State
-from exo.shared.types.tasks import ChatCompletionTask, TaskId, TaskStatus
-from exo.shared.types.worker.common import InstanceId
+from exo.shared.types.tasks import (
+    ChatCompletion as ChatCompletionTask,
+)
+from exo.shared.types.tasks import (
+    TaskId,
+    TaskStatus,
+)
+from exo.shared.types.worker.instances import InstanceId
 from exo.utils.channels import Receiver, Sender, channel
 from exo.utils.event_buffer import MultiSourceBuffer
 
diff --git a/src/exo/master/placement.py b/src/exo/master/placement.py
index 4a39ef5a..fb49666f 100644
--- a/src/exo/master/placement.py
+++ b/src/exo/master/placement.py
@@ -20,8 +20,13 @@ from exo.shared.types.common import Host
 from exo.shared.types.events import Event, InstanceCreated, InstanceDeleted
 from exo.shared.types.memory import Memory
 from exo.shared.types.topology import NodeInfo
-from exo.shared.types.worker.common import InstanceId
-from exo.shared.types.worker.instances import Instance, InstanceStatus
+from exo.shared.types.worker.instances import (
+    Instance,
+    InstanceId,
+    InstanceMeta,
+    MlxIbvInstance,
+    MlxRingInstance,
+)
 
 
 def random_ephemeral_port() -> int:
@@ -50,7 +55,6 @@ def get_instance_placements_after_create(
         raise ValueError("No cycles found with sufficient memory")
 
     smallest_cycles = get_smallest_cycles(cycles_with_sufficient_memory)
-    selected_cycle = None
 
     smallest_tb_cycles = [
         cycle
@@ -82,7 +86,7 @@ def get_instance_placements_after_create(
     )
 
     shard_assignments = get_shard_assignments(
-        command.model_meta, selected_cycle, command.strategy
+        command.model_meta, selected_cycle, command.sharding
     )
 
     cycle_digraph: Topology = topology.get_subgraph_from_nodes(selected_cycle)
@@ -90,36 +94,36 @@ def get_instance_placements_after_create(
     instance_id = InstanceId()
     target_instances = dict(deepcopy(current_instances))
 
-    if command.strategy in ("tensor_rdma", "pipeline_rdma"):
-        mlx_ibv_devices = get_mlx_ibv_devices_matrix(
-            selected_cycle,
-            cycle_digraph,
-        )
-        mlx_ibv_coordinator = get_mlx_ibv_coordinator(
-            selected_cycle,
-            coordinator_port=random_ephemeral_port(),
-        )
-        target_instances[instance_id] = Instance(
-            instance_id=instance_id,
-            instance_type=InstanceStatus.Active,
-            shard_assignments=shard_assignments,
-            mlx_ibv_devices=mlx_ibv_devices,
-            mlx_ibv_coordinator=mlx_ibv_coordinator,
-        )
-    else:
-        hosts: list[Host] = get_hosts_from_subgraph(cycle_digraph)
-        target_instances[instance_id] = Instance(
-            instance_id=instance_id,
-            instance_type=InstanceStatus.Active,
-            shard_assignments=shard_assignments,
-            hosts=[
-                Host(
-                    ip=host.ip,
-                    port=random_ephemeral_port(),
-                )
-                for host in hosts
-            ],
-        )
+    # TODO: Single node instances
+    match command.instance_meta:
+        case InstanceMeta.MlxIbv:
+            mlx_ibv_devices = get_mlx_ibv_devices_matrix(
+                selected_cycle,
+                cycle_digraph,
+            )
+            mlx_ibv_coordinator = get_mlx_ibv_coordinator(
+                selected_cycle,
+                coordinator_port=random_ephemeral_port(),
+            )
+            target_instances[instance_id] = MlxIbvInstance(
+                instance_id=instance_id,
+                shard_assignments=shard_assignments,
+                ibv_devices=mlx_ibv_devices,
+                ibv_coordinator=mlx_ibv_coordinator,
+            )
+        case InstanceMeta.MlxRing:
+            hosts: list[Host] = get_hosts_from_subgraph(cycle_digraph)
+            target_instances[instance_id] = MlxRingInstance(
+                instance_id=instance_id,
+                shard_assignments=shard_assignments,
+                hosts=[
+                    Host(
+                        ip=host.ip,
+                        port=random_ephemeral_port(),
+                    )
+                    for host in hosts
+                ],
+            )
 
     return target_instances
 
diff --git a/src/exo/master/placement_utils.py b/src/exo/master/placement_utils.py
index bd0a9073..1a4e7011 100644
--- a/src/exo/master/placement_utils.py
+++ b/src/exo/master/placement_utils.py
@@ -10,11 +10,10 @@ from exo.shared.types.memory import Memory
 from exo.shared.types.models import ModelMetadata
 from exo.shared.types.profiling import NodePerformanceProfile
 from exo.shared.types.topology import NodeInfo
-from exo.shared.types.worker.common import RunnerId
-from exo.shared.types.worker.parallelisation_strategy import ParallelisationStrategyType
-from exo.shared.types.worker.runners import ShardAssignments
+from exo.shared.types.worker.runners import RunnerId, ShardAssignments
 from exo.shared.types.worker.shards import (
     PipelineShardMetadata,
+    Sharding,
     ShardMetadata,
     TensorShardMetadata,
 )
@@ -53,7 +52,6 @@ def get_smallest_cycles(cycles: list[list[NodeInfo]]) -> list[list[NodeInfo]]:
 def get_shard_assignments_for_pipeline_parallel(
     model_meta: ModelMetadata,
     selected_cycle: list[NodeInfo],
-    parallelisation_strategy: ParallelisationStrategyType,
 ):
     if not narrow_all_nodes(selected_cycle):
         raise ValueError("All nodes must have profiles to create shard assignments")
@@ -90,7 +88,6 @@ def get_shard_assignments_for_pipeline_parallel(
             start_layer=layers_assigned,
             end_layer=layers_assigned + node_layers,
             n_layers=total_layers,
-            strategy=parallelisation_strategy,
         )
 
         runner_to_shard[runner_id] = shard
@@ -109,7 +106,6 @@ def get_shard_assignments_for_pipeline_parallel(
 def get_shard_assignments_for_tensor_parallel(
     model_meta: ModelMetadata,
     selected_cycle: list[NodeInfo],
-    parallelisation_strategy: ParallelisationStrategyType,
 ):
     if not narrow_all_nodes(selected_cycle):
         raise ValueError("All nodes must have profiles to create shard assignments")
@@ -127,7 +123,6 @@ def get_shard_assignments_for_tensor_parallel(
             start_layer=0,
             end_layer=total_layers,
             n_layers=total_layers,
-            strategy=parallelisation_strategy,
         )
 
         runner_id = RunnerId()
@@ -147,38 +142,18 @@ def get_shard_assignments_for_tensor_parallel(
 def get_shard_assignments(
     model_meta: ModelMetadata,
     selected_cycle: list[NodeInfo],
-    parallelisation_strategy: ParallelisationStrategyType,
+    sharding: Sharding,
 ) -> ShardAssignments:
-    match parallelisation_strategy:
-        case "auto":
+    match sharding:
+        case Sharding.Pipeline:
             return get_shard_assignments_for_pipeline_parallel(
                 model_meta=model_meta,
                 selected_cycle=selected_cycle,
-                parallelisation_strategy=parallelisation_strategy,
             )
-        case "pipeline":
-            return get_shard_assignments_for_pipeline_parallel(
-                model_meta=model_meta,
-                selected_cycle=selected_cycle,
-                parallelisation_strategy=parallelisation_strategy,
-            )
-        case "pipeline_rdma":
-            return get_shard_assignments_for_pipeline_parallel(
-                model_meta=model_meta,
-                selected_cycle=selected_cycle,
-                parallelisation_strategy=parallelisation_strategy,
-            )
-        case "tensor":
+        case Sharding.Tensor:
             return get_shard_assignments_for_tensor_parallel(
                 model_meta=model_meta,
                 selected_cycle=selected_cycle,
-                parallelisation_strategy=parallelisation_strategy,
-            )
-        case "tensor_rdma":
-            return get_shard_assignments_for_tensor_parallel(
-                model_meta=model_meta,
-                selected_cycle=selected_cycle,
-                parallelisation_strategy=parallelisation_strategy,
             )
 
 
@@ -300,17 +275,12 @@ def _find_interface_name_for_ip(
 def get_mlx_ibv_coordinator(
     selected_cycle: list[NodeInfo],
     coordinator_port: int,
-) -> str | None:
+) -> str:
     """Get the coordinator address for MLX IBV (rank 0 device).
 
     Selects a non-thunderbolt IP address from rank 0 node as a heuristic for
     ethernet accessibility. Returns address in format "X.X.X.X:PORT".
     """
-
-    if len(selected_cycle) == 0:
-        logger.warning("No nodes in selected cycle, cannot determine coordinator")
-        return None
-
     rank_0_node = selected_cycle[0]
     logger.info(f"Selecting coordinator from rank 0 node: {rank_0_node.node_id}")
     assert rank_0_node.node_profile is not None
diff --git a/src/exo/shared/apply.py b/src/exo/shared/apply.py
index 08150783..b30512af 100644
--- a/src/exo/shared/apply.py
+++ b/src/exo/shared/apply.py
@@ -1,5 +1,5 @@
 import copy
-from typing import Mapping
+from collections.abc import Mapping, Sequence
 
 from loguru import logger
 
@@ -8,44 +8,43 @@ from exo.shared.types.events import (
     ChunkGenerated,
     Event,
     IndexedEvent,
-    InstanceActivated,
     InstanceCreated,
-    InstanceDeactivated,
     InstanceDeleted,
+    NodeDownloadProgress,
     NodeMemoryMeasured,
     NodePerformanceMeasured,
     RunnerDeleted,
     RunnerStatusUpdated,
+    TaskAcknowledged,
     TaskCreated,
     TaskDeleted,
     TaskFailed,
-    TaskStateUpdated,
+    TaskStatusUpdated,
     TestEvent,
     TopologyEdgeCreated,
     TopologyEdgeDeleted,
     TopologyNodeCreated,
-    WorkerStatusUpdated,
 )
 from exo.shared.types.profiling import NodePerformanceProfile, SystemPerformanceProfile
 from exo.shared.types.state import State
 from exo.shared.types.tasks import Task, TaskId, TaskStatus
 from exo.shared.types.topology import NodeInfo
-from exo.shared.types.worker.common import RunnerId, WorkerStatus
-from exo.shared.types.worker.instances import Instance, InstanceId, InstanceStatus
-from exo.shared.types.worker.runners import RunnerStatus
+from exo.shared.types.worker.instances import Instance, InstanceId
+from exo.shared.types.worker.downloads import DownloadProgress
+from exo.shared.types.worker.runners import RunnerId, RunnerStatus
 
 
 def event_apply(event: Event, state: State) -> State:
     """Apply an event to state."""
     match event:
-        case TestEvent() | ChunkGenerated():
+        case (
+            TestEvent() | ChunkGenerated() | TaskAcknowledged()
+        ):  # TaskAcknowledged should never be sent by a worker but i dont mind if it just gets ignored
             return state
-        case InstanceActivated():
-            return apply_instance_activated(event, state)
+        case NodeDownloadProgress():
+            return apply_node_download_progress(event, state)
         case InstanceCreated():
             return apply_instance_created(event, state)
-        case InstanceDeactivated():
-            return apply_instance_deactivated(event, state)
         case InstanceDeleted():
             return apply_instance_deleted(event, state)
         case NodePerformanceMeasured():
@@ -62,10 +61,8 @@ def event_apply(event: Event, state: State) -> State:
             return apply_task_deleted(event, state)
         case TaskFailed():
             return apply_task_failed(event, state)
-        case TaskStateUpdated():
-            return apply_task_state_updated(event, state)
-        case WorkerStatusUpdated():
-            return apply_worker_status_updated(event, state)
+        case TaskStatusUpdated():
+            return apply_task_status_updated(event, state)
         case TopologyNodeCreated():
             return apply_topology_node_created(event, state)
         case TopologyEdgeCreated():
@@ -85,6 +82,22 @@ def apply(state: State, event: IndexedEvent) -> State:
     return new_state.model_copy(update={"last_event_applied_idx": event.idx})
 
 
+def apply_node_download_progress(event: NodeDownloadProgress, state: State) -> State:
+    new_node_downloads: Sequence[DownloadProgress] = [
+        event.download_progress
+        if dp.shard_metadata == event.download_progress.shard_metadata
+        else dp
+        for dp in state.downloads.get(
+            event.download_progress.node_id, [event.download_progress]
+        )
+    ]
+    new_downloads: Mapping[NodeId, Sequence[DownloadProgress]] = {
+        **state.downloads,
+        event.download_progress.node_id: new_node_downloads,
+    }
+    return state.model_copy(update={"downloads": new_downloads})
+
+
 def apply_task_created(event: TaskCreated, state: State) -> State:
     new_tasks: Mapping[TaskId, Task] = {**state.tasks, event.task_id: event.task}
     return state.model_copy(update={"tasks": new_tasks})
@@ -97,8 +110,9 @@ def apply_task_deleted(event: TaskDeleted, state: State) -> State:
     return state.model_copy(update={"tasks": new_tasks})
 
 
-def apply_task_state_updated(event: TaskStateUpdated, state: State) -> State:
+def apply_task_status_updated(event: TaskStatusUpdated, state: State) -> State:
     if event.task_id not in state.tasks:
+        # maybe should raise
         return state
 
     update: dict[str, TaskStatus | None] = {
@@ -115,6 +129,7 @@ def apply_task_state_updated(event: TaskStateUpdated, state: State) -> State:
 
 def apply_task_failed(event: TaskFailed, state: State) -> State:
     if event.task_id not in state.tasks:
+        # maybe should raise
         return state
 
     updated_task = state.tasks[event.task_id].model_copy(
@@ -133,34 +148,6 @@ def apply_instance_created(event: InstanceCreated, state: State) -> State:
     return state.model_copy(update={"instances": new_instances})
 
 
-def apply_instance_activated(event: InstanceActivated, state: State) -> State:
-    if event.instance_id not in state.instances:
-        return state
-
-    updated_instance = state.instances[event.instance_id].model_copy(
-        update={"instance_type": InstanceStatus.Active}
-    )
-    new_instances: Mapping[InstanceId, Instance] = {
-        **state.instances,
-        event.instance_id: updated_instance,
-    }
-    return state.model_copy(update={"instances": new_instances})
-
-
-def apply_instance_deactivated(event: InstanceDeactivated, state: State) -> State:
-    if event.instance_id not in state.instances:
-        return state
-
-    updated_instance = state.instances[event.instance_id].model_copy(
-        update={"instance_type": InstanceStatus.Inactive}
-    )
-    new_instances: Mapping[InstanceId, Instance] = {
-        **state.instances,
-        event.instance_id: updated_instance,
-    }
-    return state.model_copy(update={"instances": new_instances})
-
-
 def apply_instance_deleted(event: InstanceDeleted, state: State) -> State:
     new_instances: Mapping[InstanceId, Instance] = {
         iid: inst for iid, inst in state.instances.items() if iid != event.instance_id
@@ -177,6 +164,9 @@ def apply_runner_status_updated(event: RunnerStatusUpdated, state: State) -> Sta
 
 
 def apply_runner_deleted(event: RunnerDeleted, state: State) -> State:
+    assert event.runner_id in state.runners, (
+        "RunnerDeleted before any RunnerStatusUpdated events"
+    )
     new_runners: Mapping[RunnerId, RunnerStatus] = {
         rid: rs for rid, rs in state.runners.items() if rid != event.runner_id
     }
@@ -245,14 +235,6 @@ def apply_node_memory_measured(event: NodeMemoryMeasured, state: State) -> State
     )
 
 
-def apply_worker_status_updated(event: WorkerStatusUpdated, state: State) -> State:
-    new_node_status: Mapping[NodeId, WorkerStatus] = {
-        **state.node_status,
-        event.node_id: event.node_state,
-    }
-    return state.model_copy(update={"node_status": new_node_status})
-
-
 def apply_topology_node_created(event: TopologyNodeCreated, state: State) -> State:
     topology = copy.copy(state.topology)
     topology.add_node(NodeInfo(node_id=event.node_id))
diff --git a/src/exo/shared/global_conn.py b/src/exo/shared/global_conn.py
deleted file mode 100644
index 01bfc203..00000000
--- a/src/exo/shared/global_conn.py
+++ /dev/null
@@ -1,67 +0,0 @@
-# src/exo/shared/global_conn.py
-
-import asyncio
-import threading
-from multiprocessing.connection import Connection
-
-from exo.shared.types.worker.commands_runner import (
-    RunnerMessage,
-    RunnerResponse,
-)
-
-
-class AsyncConnection[SendT, RecvT]:
-    """
-    Async/sync wrapper around multiprocessing.Connection with thread-safe send.
-    Use:
-      - await send(...) from asyncio code
-      - send_sync(...) from executor/background threads
-    """
-
-    def __init__(self, conn: Connection):
-        self._conn = conn
-        self._send_lock = threading.Lock()
-        self._recv_lock = threading.Lock()
-
-    # ---- sending ----
-    async def send(self, obj: SendT) -> None:
-        loop = asyncio.get_running_loop()
-        await loop.run_in_executor(None, self._send_blocking, obj)
-
-    def send_sync(self, obj: SendT) -> None:
-        self._send_blocking(obj)
-
-    def _send_blocking(self, obj: SendT) -> None:
-        # Single critical section for the whole pickle frame
-        with self._send_lock:
-            self._conn.send(obj)
-
-    # ---- receiving ----
-    async def recv(self) -> RecvT:
-        loop = asyncio.get_running_loop()
-        return await loop.run_in_executor(None, self._recv_blocking)
-
-    def _recv_blocking(self) -> RecvT:
-        # Not strictly needed in your parent, but safe if misused elsewhere
-        with self._recv_lock:
-            return self._conn.recv()  # type: ignore[no-any-return]
-
-    async def poll(self, timeout: float | None = None) -> bool:
-        return await asyncio.to_thread(self._conn.poll, timeout)
-
-    def close(self) -> None:
-        self._conn.close()
-
-
-_conn: AsyncConnection[RunnerResponse, RunnerMessage] | None = None
-
-
-def set_conn(c: AsyncConnection[RunnerResponse, RunnerMessage]) -> None:
-    global _conn
-    _conn = c
-
-
-def get_conn() -> AsyncConnection[RunnerResponse, RunnerMessage]:
-    if _conn is None:
-        raise RuntimeError("Global conn has not been set yet")
-    return _conn
diff --git a/src/exo/shared/logging.py b/src/exo/shared/logging.py
index 60705bf6..66ba1700 100644
--- a/src/exo/shared/logging.py
+++ b/src/exo/shared/logging.py
@@ -4,11 +4,11 @@ from pathlib import Path
 from loguru import logger
 
 
-def logger_setup(log_file: Path, verbosity: int = 0):
+def logger_setup(log_file: Path | None, verbosity: int = 0):
     """Set up logging for this process - formatting, file handles, verbosity and output"""
     logger.remove()
     if verbosity == 0:
-        _ = logger.add(  # type: ignore
+        logger.add(
             sys.__stderr__,  # type: ignore
             format="[ {time:hh:mm:ss.SSSSA} | <level>{level: <8}</level>] <level>{message}</level>",
             level="INFO",
@@ -16,19 +16,22 @@ def logger_setup(log_file: Path, verbosity: int = 0):
             enqueue=True,
         )
     else:
-        _ = logger.add(  # type: ignore
+        logger.add(
             sys.__stderr__,  # type: ignore
             format="[ {time:HH:mm:ss.SSS} | <level>{level: <8}</level> | {name}:{function}:{line} ] <level>{message}</level>",
             level="DEBUG",
             colorize=True,
             enqueue=True,
         )
-    _ = logger.add(
-        log_file,
-        format="[ {time:YYYY-MM-DD HH:mm:ss.SSS} | {level: <8} | {name}:{function}:{line} ] {message}",
-        level="INFO",
-        enqueue=True,
-    )
+    if log_file:
+        logger.add(
+            log_file,
+            format="[ {time:YYYY-MM-DD HH:mm:ss.SSS} | {level: <8} | {name}:{function}:{line} ] {message}",
+            level="INFO",
+            colorize=False,
+            enqueue=True,
+            rotation="1 week",
+        )
 
 
 def logger_cleanup():
diff --git a/src/exo/shared/types/api.py b/src/exo/shared/types/api.py
index e5437c4e..131fc7e2 100644
--- a/src/exo/shared/types/api.py
+++ b/src/exo/shared/types/api.py
@@ -6,8 +6,8 @@ from pydantic import BaseModel, Field
 from exo.shared.openai_compat import FinishReason
 from exo.shared.types.common import CommandId
 from exo.shared.types.models import ModelMetadata
-from exo.shared.types.worker.instances import InstanceId
-from exo.shared.types.worker.parallelisation_strategy import ParallelisationStrategyType
+from exo.shared.types.worker.instances import InstanceId, InstanceMeta
+from exo.shared.types.worker.shards import Sharding
 
 
 class ModelListModel(BaseModel):
@@ -124,7 +124,9 @@ class ChatCompletionTaskParams(BaseModel):
 class CreateInstanceTaskParams(BaseModel):
     # TODO: in future the user could specify a specific Instance, not just a model_id
     model_id: str
-    strategy: ParallelisationStrategyType = "auto"
+    sharding: Sharding = Sharding.Pipeline
+    # TODO: fix
+    instance_meta: InstanceMeta = InstanceMeta.MlxRing
 
 
 class DeleteInstanceTaskParams(BaseModel):
diff --git a/src/exo/shared/types/chunks.py b/src/exo/shared/types/chunks.py
index f74c901a..990416c0 100644
--- a/src/exo/shared/types/chunks.py
+++ b/src/exo/shared/types/chunks.py
@@ -1,7 +1,6 @@
 from enum import Enum
 
 from exo.shared.openai_compat import FinishReason
-from exo.shared.types.common import CommandId
 from exo.shared.types.models import ModelId
 from exo.utils.pydantic_ext import TaggedModel
 
@@ -12,7 +11,6 @@ class ChunkType(str, Enum):
 
 
 class BaseChunk(TaggedModel):
-    command_id: CommandId
     idx: int
     model: ModelId
 
diff --git a/src/exo/shared/types/commands.py b/src/exo/shared/types/commands.py
index d746cca2..979d42bd 100644
--- a/src/exo/shared/types/commands.py
+++ b/src/exo/shared/types/commands.py
@@ -3,8 +3,8 @@ from pydantic import Field
 from exo.shared.types.api import ChatCompletionTaskParams
 from exo.shared.types.common import CommandId, NodeId
 from exo.shared.types.models import ModelMetadata
-from exo.shared.types.worker.common import InstanceId
-from exo.shared.types.worker.parallelisation_strategy import ParallelisationStrategyType
+from exo.shared.types.worker.instances import InstanceId, InstanceMeta
+from exo.shared.types.worker.shards import Sharding
 from exo.utils.pydantic_ext import CamelCaseModel, TaggedModel
 
 
@@ -16,6 +16,8 @@ class BaseCommand(TaggedModel):
 class TestCommand(BaseCommand):
     pass
 
+class KillCommand(BaseCommand):
+    pass
 
 class ChatCompletion(BaseCommand):
     request_params: ChatCompletionTaskParams
@@ -23,7 +25,8 @@ class ChatCompletion(BaseCommand):
 
 class CreateInstance(BaseCommand):
     model_meta: ModelMetadata
-    strategy: ParallelisationStrategyType
+    sharding: Sharding
+    instance_meta: InstanceMeta
 
 
 class SpinUpInstance(BaseCommand):
@@ -44,6 +47,7 @@ class RequestEventLog(BaseCommand):
 
 Command = (
     TestCommand
+    | KillCommand
     | RequestEventLog
     | ChatCompletion
     | CreateInstance
diff --git a/src/exo/shared/types/events.py b/src/exo/shared/types/events.py
index c0083809..ccc88185 100644
--- a/src/exo/shared/types/events.py
+++ b/src/exo/shared/types/events.py
@@ -1,15 +1,14 @@
 from datetime import datetime
-from enum import Enum
 
 from pydantic import Field
 
 from exo.shared.topology import Connection, NodePerformanceProfile
-from exo.shared.types.chunks import CommandId, GenerationChunk
-from exo.shared.types.common import Id, NodeId, SessionId
+from exo.shared.types.chunks import GenerationChunk
+from exo.shared.types.common import CommandId, Id, NodeId, SessionId
 from exo.shared.types.profiling import MemoryPerformanceProfile
 from exo.shared.types.tasks import Task, TaskId, TaskStatus
-from exo.shared.types.worker.common import InstanceId, WorkerStatus
-from exo.shared.types.worker.instances import Instance
+from exo.shared.types.worker.downloads import DownloadProgress
+from exo.shared.types.worker.instances import Instance, InstanceId
 from exo.shared.types.worker.runners import RunnerId, RunnerStatus
 from exo.utils.pydantic_ext import CamelCaseModel, TaggedModel
 
@@ -20,45 +19,6 @@ class EventId(Id):
     """
 
 
-class EventType(str, Enum):
-    """
-    Here are all the unique kinds of events that can be sent over the network.
-    """
-
-    # Test Events, strictly for mocks and tests.
-    TestEvent = "TestEvent"
-
-    # Task Events
-    TaskCreated = "TaskCreated"
-    TaskStateUpdated = "TaskStateUpdated"
-    TaskFailed = "TaskFailed"
-    TaskDeleted = "TaskDeleted"
-
-    # Streaming Events
-    ChunkGenerated = "ChunkGenerated"
-
-    # Instance Events
-    InstanceCreated = "InstanceCreated"
-    InstanceDeleted = "InstanceDeleted"
-    InstanceActivated = "InstanceActivated"
-    InstanceDeactivated = "InstanceDeactivated"
-    InstanceReplacedAtomically = "InstanceReplacedAtomically"
-
-    # Runner Status Events
-    RunnerStatusUpdated = "RunnerStatusUpdated"
-    RunnerDeleted = "RunnerDeleted"
-
-    # Node Performance Events
-    WorkerStatusUpdated = "WorkerStatusUpdated"
-    NodePerformanceMeasured = "NodePerformanceMeasured"
-    NodeMemoryMeasured = "NodeMemoryMeasured"
-
-    # Topology Events
-    TopologyNodeCreated = "TopologyNodeCreated"
-    TopologyEdgeCreated = "TopologyEdgeCreated"
-    TopologyEdgeDeleted = "TopologyEdgeDeleted"
-
-
 class BaseEvent(TaggedModel):
     event_id: EventId = Field(default_factory=EventId)
     # Internal, for debugging. Please don't rely on this field for anything!
@@ -74,11 +34,15 @@ class TaskCreated(BaseEvent):
     task: Task
 
 
+class TaskAcknowledged(BaseEvent):
+    task_id: TaskId
+
+
 class TaskDeleted(BaseEvent):
     task_id: TaskId
 
 
-class TaskStateUpdated(BaseEvent):
+class TaskStatusUpdated(BaseEvent):
     task_id: TaskId
     task_status: TaskStatus
 
@@ -93,14 +57,6 @@ class InstanceCreated(BaseEvent):
     instance: Instance
 
 
-class InstanceActivated(BaseEvent):
-    instance_id: InstanceId
-
-
-class InstanceDeactivated(BaseEvent):
-    instance_id: InstanceId
-
-
 class InstanceDeleted(BaseEvent):
     instance_id: InstanceId
 
@@ -119,14 +75,13 @@ class NodePerformanceMeasured(BaseEvent):
     node_profile: NodePerformanceProfile
 
 
-class NodeMemoryMeasured(BaseEvent):
-    node_id: NodeId
-    memory: MemoryPerformanceProfile
+class NodeDownloadProgress(BaseEvent):
+    download_progress: DownloadProgress
 
 
-class WorkerStatusUpdated(BaseEvent):
+class NodeMemoryMeasured(BaseEvent):
     node_id: NodeId
-    node_state: WorkerStatus
+    memory: MemoryPerformanceProfile
 
 
 class ChunkGenerated(BaseEvent):
@@ -149,18 +104,17 @@ class TopologyEdgeDeleted(BaseEvent):
 Event = (
     TestEvent
     | TaskCreated
-    | TaskStateUpdated
+    | TaskStatusUpdated
     | TaskFailed
     | TaskDeleted
+    | TaskAcknowledged
     | InstanceCreated
-    | InstanceActivated
-    | InstanceDeactivated
     | InstanceDeleted
     | RunnerStatusUpdated
     | RunnerDeleted
     | NodePerformanceMeasured
     | NodeMemoryMeasured
-    | WorkerStatusUpdated
+    | NodeDownloadProgress
     | ChunkGenerated
     | TopologyNodeCreated
     | TopologyEdgeCreated
diff --git a/src/exo/shared/types/memory.py b/src/exo/shared/types/memory.py
index 21cd1534..562c3c87 100644
--- a/src/exo/shared/types/memory.py
+++ b/src/exo/shared/types/memory.py
@@ -47,6 +47,11 @@ class Memory(CamelCaseModel):
         """Construct a new Memory object from a number of megabytes"""
         return cls(in_bytes=round(val * (1024**2)))
 
+    @property
+    def in_gb(self) -> float:
+        """The approximate gigabytes this memory represents."""
+        return self.in_bytes / (1024**3)
+
     def __add__(self, other: "Memory") -> "Memory":
         return Memory.from_bytes(self.in_bytes + other.in_bytes)
 
diff --git a/src/exo/shared/types/state.py b/src/exo/shared/types/state.py
index fcdd08a6..3cd3a256 100644
--- a/src/exo/shared/types/state.py
+++ b/src/exo/shared/types/state.py
@@ -2,14 +2,15 @@ from collections.abc import Mapping, Sequence
 from typing import Any, cast
 
 from pydantic import ConfigDict, Field, field_serializer, field_validator
+from pydantic.alias_generators import to_camel
 
 from exo.shared.topology import Topology, TopologySnapshot
 from exo.shared.types.common import NodeId
 from exo.shared.types.profiling import NodePerformanceProfile
 from exo.shared.types.tasks import Task, TaskId
-from exo.shared.types.worker.common import InstanceId, WorkerStatus
-from exo.shared.types.worker.instances import Instance
+from exo.shared.types.worker.instances import Instance, InstanceId
 from exo.shared.types.worker.runners import RunnerId, RunnerStatus
+from exo.shared.types.worker.downloads import DownloadProgress
 from exo.utils.pydantic_ext import CamelCaseModel
 
 
@@ -22,15 +23,19 @@ class State(CamelCaseModel):
     """
 
     model_config = ConfigDict(
+        alias_generator=to_camel,
+        validate_by_name=True,
+        extra="forbid",
+        # I want to reenable this ASAP, but it's causing an issue with TaskStatus
+        strict=True,
         arbitrary_types_allowed=True,
     )
-    node_status: Mapping[NodeId, WorkerStatus] = {}
     instances: Mapping[InstanceId, Instance] = {}
     runners: Mapping[RunnerId, RunnerStatus] = {}
+    downloads: Mapping[NodeId, Sequence[DownloadProgress]] = {}
     tasks: Mapping[TaskId, Task] = {}
     node_profiles: Mapping[NodeId, NodePerformanceProfile] = {}
     topology: Topology = Topology()
-    history: Sequence[Topology] = []
     last_event_applied_idx: int = Field(default=-1, ge=-1)
 
     @field_serializer("topology", mode="plain")
diff --git a/src/exo/shared/types/tasks.py b/src/exo/shared/types/tasks.py
index 0e38d5dc..40fb1611 100644
--- a/src/exo/shared/types/tasks.py
+++ b/src/exo/shared/types/tasks.py
@@ -4,7 +4,9 @@ from pydantic import Field
 
 from exo.shared.types.api import ChatCompletionTaskParams
 from exo.shared.types.common import CommandId, Id
-from exo.shared.types.worker.common import InstanceId
+from exo.shared.types.worker.instances import BoundInstance, InstanceId
+from exo.shared.types.worker.runners import RunnerId
+from exo.shared.types.worker.shards import ShardMetadata
 from exo.utils.pydantic_ext import TaggedModel
 
 
@@ -19,15 +21,40 @@ class TaskStatus(str, Enum):
     Failed = "Failed"
 
 
-class ChatCompletionTask(TaggedModel):
-    task_id: TaskId
-    command_id: CommandId
+class BaseTask(TaggedModel):
+    task_id: TaskId = Field(default_factory=TaskId)
+    task_status: TaskStatus = Field(default=TaskStatus.Pending)
     instance_id: InstanceId
-    task_status: TaskStatus
+
+
+class CreateRunner(BaseTask):  # emitted by Worker
+    bound_instance: BoundInstance
+
+
+class DownloadModel(BaseTask):  # emitted by Worker
+    shard_metadata: ShardMetadata
+
+
+class LoadModel(BaseTask):  # emitted by Worker
+    pass
+
+
+class StartWarmup(BaseTask):  # emitted by Worker
+    pass
+
+
+class ChatCompletion(BaseTask):  # emitted by Master
+    command_id: CommandId
     task_params: ChatCompletionTaskParams
 
     error_type: str | None = Field(default=None)
     error_message: str | None = Field(default=None)
 
 
-Task = ChatCompletionTask
+class Shutdown(BaseTask):  # emitted by Worker
+    runner_id: RunnerId
+
+
+Task = (
+    CreateRunner | DownloadModel | LoadModel | StartWarmup | ChatCompletion | Shutdown
+)
diff --git a/src/exo/shared/types/worker/commands_runner.py b/src/exo/shared/types/worker/commands_runner.py
index 0a873f8a..8878937f 100644
--- a/src/exo/shared/types/worker/commands_runner.py
+++ b/src/exo/shared/types/worker/commands_runner.py
@@ -1,33 +1,7 @@
 from exo.shared.openai_compat import FinishReason
-from exo.shared.types.common import Host
-from exo.shared.types.tasks import ChatCompletionTaskParams
-from exo.shared.types.worker.shards import ShardMetadata
 from exo.utils.pydantic_ext import TaggedModel
 
 
-class BaseRunnerMessage(TaggedModel):
-    pass
-
-
-class SetupMessage(BaseRunnerMessage):
-    model_shard_meta: ShardMetadata
-    hosts: list[Host] | None = None
-    mlx_ibv_devices: list[list[str | None]] | None = None
-    mlx_ibv_coordinator: str | None = None
-
-
-# TODO: We probably want a general task message that can take any task type. Can be fixed later.
-class ChatTaskMessage(BaseRunnerMessage):
-    task_data: ChatCompletionTaskParams
-
-
-class ExitMessage(BaseRunnerMessage):
-    pass
-
-
-RunnerMessage = SetupMessage | ChatTaskMessage | ExitMessage
-
-
 class BaseRunnerResponse(TaggedModel):
     pass
 
diff --git a/src/exo/shared/types/worker/common.py b/src/exo/shared/types/worker/common.py
index 6dd29380..8b137891 100644
--- a/src/exo/shared/types/worker/common.py
+++ b/src/exo/shared/types/worker/common.py
@@ -1,26 +1 @@
-from enum import Enum
 
-from exo.shared.types.common import Id
-
-
-class InstanceId(Id):
-    pass
-
-
-class RunnerId(Id):
-    pass
-
-
-class WorkerStatus(str, Enum):
-    Idle = "Idle"
-    Running = "Running"
-
-
-class RunnerError(Exception):
-    """Exception raised when the runner process encounters an error."""
-
-    def __init__(self, error_type: str, error_message: str, traceback: str):
-        self.error_type = error_type
-        self.error_message = error_message
-        self.traceback = traceback
-        super().__init__(f"{error_type}: {error_message}. Traceback: {traceback}")
diff --git a/src/exo/shared/types/worker/communication.py b/src/exo/shared/types/worker/communication.py
deleted file mode 100644
index 7643af88..00000000
--- a/src/exo/shared/types/worker/communication.py
+++ /dev/null
@@ -1,40 +0,0 @@
-import asyncio
-import traceback
-
-from loguru import logger
-
-from exo.shared.global_conn import AsyncConnection, get_conn
-from exo.shared.types.worker.commands_runner import (
-    ErrorResponse,
-    PrintResponse,
-    RunnerMessage,
-    RunnerResponse,
-)
-
-### Utils - Runner Prints
-
-
-def runner_print(text: str) -> None:
-    obj = PrintResponse(
-        text=text,
-    )
-
-    conn: AsyncConnection[RunnerResponse, RunnerMessage] = get_conn()
-    conn.send_sync(obj)
-
-
-def runner_write_error(error: Exception) -> None:
-    error_response: ErrorResponse = ErrorResponse(
-        error_type=type(error).__name__,
-        error_message=str(error),
-        traceback=traceback.format_exc(),
-    )
-
-    conn = get_conn()
-    asyncio.create_task(conn.send(error_response))
-    logger.opt(exception=error).exception("Critical Runner error")
-
-
-## TODO: To make this cleaner, it seems like we should have only one writer.
-# This is fine in runner_supervisor but there's a risk in runner.py that we overlap things
-# We can guarantee this by enqueueing messages and have a writing thread.
diff --git a/src/exo/shared/types/worker/downloads.py b/src/exo/shared/types/worker/downloads.py
index 96c31b7d..73255f62 100644
--- a/src/exo/shared/types/worker/downloads.py
+++ b/src/exo/shared/types/worker/downloads.py
@@ -1,5 +1,6 @@
 from exo.shared.types.common import NodeId
 from exo.shared.types.memory import Memory
+from exo.shared.types.worker.shards import ShardMetadata
 from exo.utils.pydantic_ext import CamelCaseModel, TaggedModel
 
 
@@ -19,6 +20,7 @@ class DownloadProgressData(CamelCaseModel):
 
 class BaseDownloadProgress(TaggedModel):
     node_id: NodeId
+    shard_metadata: ShardMetadata
 
 
 class DownloadPending(BaseDownloadProgress):
diff --git a/src/exo/shared/types/worker/instances.py b/src/exo/shared/types/worker/instances.py
index 6973a48f..9230001f 100644
--- a/src/exo/shared/types/worker/instances.py
+++ b/src/exo/shared/types/worker/instances.py
@@ -1,22 +1,56 @@
 from enum import Enum
 
-from exo.shared.types.common import Host
-from exo.shared.types.worker.common import InstanceId
-from exo.shared.types.worker.runners import (
-    ShardAssignments,
-)
-from exo.utils.pydantic_ext import CamelCaseModel
+from pydantic import model_validator
 
+from exo.shared.types.common import Host, Id
+from exo.shared.types.worker.runners import RunnerId, ShardAssignments, ShardMetadata
+from exo.utils.pydantic_ext import CamelCaseModel, TaggedModel
 
-class InstanceStatus(str, Enum):
-    Active = "Active"
-    Inactive = "Inactive"
 
+class InstanceId(Id):
+    pass
 
-class Instance(CamelCaseModel):
+
+class InstanceMeta(str, Enum):
+    MlxRing = "MlxRing"
+    MlxIbv = "MlxIbv"
+
+
+class BaseInstance(TaggedModel):
     instance_id: InstanceId
-    instance_type: InstanceStatus
     shard_assignments: ShardAssignments
-    hosts: list[Host] | None = None
-    mlx_ibv_devices: list[list[str | None]] | None = None
-    mlx_ibv_coordinator: str | None = None
+
+    def shard(self, runner_id: RunnerId) -> ShardMetadata | None:
+        return self.shard_assignments.runner_to_shard.get(runner_id, None)
+
+
+class MlxRingInstance(BaseInstance):
+    hosts: list[Host]
+
+
+class MlxIbvInstance(BaseInstance):
+    ibv_devices: list[list[str | None]]
+    ibv_coordinator: str
+
+
+# TODO: Single node instance
+Instance = MlxRingInstance | MlxIbvInstance
+
+
+class BoundInstance(CamelCaseModel):
+    instance: Instance
+    bound_runner_id: RunnerId
+
+    def bound_shard(self) -> ShardMetadata:
+        shard = self.instance.shard(self.bound_runner_id)
+        assert shard is not None
+        return shard
+
+    @model_validator(mode="after")
+    def validate_shard_exists(self) -> "BoundInstance":
+        assert (
+            self.bound_runner_id in self.instance.shard_assignments.runner_to_shard
+        ), (
+            "Bound Instance must be constructed with a runner_id that is in the instances assigned shards"
+        )
+        return self
diff --git a/src/exo/shared/types/worker/ops.py b/src/exo/shared/types/worker/ops.py
index bc53feaa..5dd98c9a 100644
--- a/src/exo/shared/types/worker/ops.py
+++ b/src/exo/shared/types/worker/ops.py
@@ -1,51 +1,34 @@
-from exo.shared.types.common import Host
-from exo.shared.types.events import InstanceId
 from exo.shared.types.tasks import Task
-from exo.shared.types.worker.common import RunnerId
-from exo.shared.types.worker.shards import ShardMetadata
+from exo.shared.types.worker.instances import BoundInstance, Instance
+from exo.shared.types.worker.runners import RunnerId
 from exo.utils.pydantic_ext import TaggedModel
 
 
 class BaseRunnerOp(TaggedModel):
-    pass
+    runner_id: RunnerId
 
 
 class AssignRunnerOp(BaseRunnerOp):
-    instance_id: InstanceId
-    runner_id: RunnerId
-    shard_metadata: ShardMetadata
-    hosts: list[Host] | None = None
-    mlx_ibv_devices: list[list[str | None]] | None = None
-    mlx_ibv_coordinator: str | None = None
+    instance: Instance
+
+    def bound_instance(self) -> BoundInstance:
+        return BoundInstance(instance=self.instance, bound_runner_id=self.runner_id)
 
 
 class UnassignRunnerOp(BaseRunnerOp):
-    runner_id: RunnerId
+    pass
 
 
 class RunnerUpOp(BaseRunnerOp):
-    runner_id: RunnerId
+    pass
 
 
 class RunnerDownOp(BaseRunnerOp):
-    runner_id: RunnerId
-
-
-class RunnerFailedOp(BaseRunnerOp):
-    runner_id: RunnerId
+    pass
 
 
 class ExecuteTaskOp(BaseRunnerOp):
-    runner_id: RunnerId
     task: Task
 
 
-# Aggregate all runner operations into a single, strictly-typed union for dispatching.
-RunnerOp = (
-    AssignRunnerOp
-    | UnassignRunnerOp
-    | RunnerUpOp
-    | RunnerDownOp
-    | RunnerFailedOp
-    | ExecuteTaskOp
-)
+RunnerOp = AssignRunnerOp | ExecuteTaskOp | UnassignRunnerOp | RunnerUpOp | RunnerDownOp
diff --git a/src/exo/shared/types/worker/parallelisation_strategy.py b/src/exo/shared/types/worker/parallelisation_strategy.py
deleted file mode 100644
index e02ba89b..00000000
--- a/src/exo/shared/types/worker/parallelisation_strategy.py
+++ /dev/null
@@ -1,13 +0,0 @@
-from typing import Literal
-
-ParallelisationStrategyType = Literal[
-    "auto",
-    "pipeline",
-    "tensor",
-    "tensor_rdma",
-    "pipeline_rdma",
-]
-
-
-def strategy_error() -> ValueError:
-    return ValueError("Unexpected strategy")
diff --git a/src/exo/shared/types/worker/runners.py b/src/exo/shared/types/worker/runners.py
index 1a36a268..dd1d7271 100644
--- a/src/exo/shared/types/worker/runners.py
+++ b/src/exo/shared/types/worker/runners.py
@@ -2,49 +2,61 @@ from collections.abc import Mapping
 
 from pydantic import model_validator
 
-from exo.shared.types.common import NodeId
+from exo.shared.types.common import Id, NodeId
 from exo.shared.types.models import ModelId
-from exo.shared.types.worker.common import RunnerId
-from exo.shared.types.worker.downloads import DownloadProgress
 from exo.shared.types.worker.shards import ShardMetadata
 from exo.utils.pydantic_ext import CamelCaseModel, TaggedModel
 
 
+class RunnerId(Id):
+    pass
+
+
+class RunnerError(Exception):
+    pass
+
+
 class BaseRunnerStatus(TaggedModel):
+    def is_running(self):
+        return isinstance(self, RunnerRunning)
+
+
+class RunnerWaitingForModel(BaseRunnerStatus):
     pass
 
 
-class DownloadingRunnerStatus(BaseRunnerStatus):
-    download_progress: DownloadProgress
+class RunnerLoading(BaseRunnerStatus):
+    pass
 
 
-class InactiveRunnerStatus(BaseRunnerStatus):
+class RunnerLoaded(BaseRunnerStatus):
     pass
 
 
-class StartingRunnerStatus(BaseRunnerStatus):
+class RunnerWarmingUp(BaseRunnerStatus):
     pass
 
 
-class LoadedRunnerStatus(BaseRunnerStatus):
+class RunnerReady(BaseRunnerStatus):
     pass
 
 
-class RunningRunnerStatus(BaseRunnerStatus):
+class RunnerRunning(BaseRunnerStatus):
     pass
 
 
-class FailedRunnerStatus(BaseRunnerStatus):
+class RunnerFailed(BaseRunnerStatus):
     error_message: str | None = None
 
 
 RunnerStatus = (
-    DownloadingRunnerStatus
-    | InactiveRunnerStatus
-    | StartingRunnerStatus
-    | LoadedRunnerStatus
-    | RunningRunnerStatus
-    | FailedRunnerStatus
+    RunnerWaitingForModel
+    | RunnerLoading
+    | RunnerLoaded
+    | RunnerWarmingUp
+    | RunnerReady
+    | RunnerRunning
+    | RunnerFailed
 )
 
 
diff --git a/src/exo/shared/types/worker/shards.py b/src/exo/shared/types/worker/shards.py
index 7270fba5..303adcc3 100644
--- a/src/exo/shared/types/worker/shards.py
+++ b/src/exo/shared/types/worker/shards.py
@@ -1,10 +1,16 @@
-from pydantic import Field
+from enum import Enum
+
+from pydantic import Field, ConfigDict
 
 from exo.shared.types.models import ModelMetadata
-from exo.shared.types.worker.parallelisation_strategy import ParallelisationStrategyType
 from exo.utils.pydantic_ext import TaggedModel
 
 
+class Sharding(str, Enum):
+    Tensor = "Tensor"
+    Pipeline = "Pipeline"
+
+
 class BaseShardMetadata(TaggedModel):
     """
     Defines a specific shard of the model that is ready to be run on a device.
@@ -24,8 +30,6 @@ class BaseShardMetadata(TaggedModel):
     end_layer: int = Field(ge=0)
     n_layers: int = Field(ge=0)
 
-    strategy: ParallelisationStrategyType = "auto"
-
     @property
     def is_first_layer(self) -> bool:
         return self.start_layer == 0
@@ -36,7 +40,14 @@ class BaseShardMetadata(TaggedModel):
 
     def __hash__(self) -> int:
         return hash(
-            (self.model_meta.model_id, self.start_layer, self.end_layer, self.n_layers)
+            (
+                self.model_meta.model_id,
+                self.start_layer,
+                self.end_layer,
+                self.n_layers,
+                self.device_rank,
+                self.world_size,
+            )
         )
 
 
@@ -48,11 +59,9 @@ class PipelineShardMetadata(BaseShardMetadata):
     where start_layer is inclusive and end_layer is exclusive.
     """
 
-    strategy: ParallelisationStrategyType = "pipeline"
-
 
 class TensorShardMetadata(BaseShardMetadata):
-    strategy: ParallelisationStrategyType = "tensor"
+    pass
 
 
 ShardMetadata = PipelineShardMetadata | TensorShardMetadata
diff --git a/src/exo/utils/channels.py b/src/exo/utils/channels.py
index 8450a664..2849f076 100644
--- a/src/exo/utils/channels.py
+++ b/src/exo/utils/channels.py
@@ -1,7 +1,18 @@
+import multiprocessing as mp
+from dataclasses import dataclass, field
 from math import inf
+from multiprocessing.synchronize import Event
+from queue import Empty, Full
+from types import TracebackType
 from typing import Self
 
-from anyio import ClosedResourceError, WouldBlock
+from anyio import (
+    CapacityLimiter,
+    ClosedResourceError,
+    EndOfStream,
+    WouldBlock,
+    to_thread,
+)
 from anyio.streams.memory import (
     MemoryObjectReceiveStream as AnyioReceiver,
 )
@@ -14,6 +25,11 @@ from anyio.streams.memory import (
 
 
 class Sender[T](AnyioSender[T]):
+    def clone(self) -> "Sender[T]":
+        if self._closed:
+            raise ClosedResourceError
+        return Sender(_state=self._state)
+
     def clone_receiver(self) -> "Receiver[T]":
         """Constructs a Receiver using a Senders shared state - similar to calling Receiver.clone() without needing the receiver"""
         if self._closed:
@@ -22,6 +38,11 @@ class Sender[T](AnyioSender[T]):
 
 
 class Receiver[T](AnyioReceiver[T]):
+    def clone(self) -> "Receiver[T]":
+        if self._closed:
+            raise ClosedResourceError
+        return Receiver(_state=self._state)
+
     def clone_sender(self) -> Sender[T]:
         """Constructs a Sender using a Receivers shared state - similar to calling Sender.clone() without needing the sender"""
         if self._closed:
@@ -52,9 +73,210 @@ class Receiver[T](AnyioReceiver[T]):
         return self
 
 
+class _MpEndOfStream:
+    pass
+
+
+MP_END_OF_STREAM = _MpEndOfStream()
+
+
+class MpState[T]:
+    def __init__(self, max_buffer_size: float):
+        if max_buffer_size == inf:
+            max_buffer_size = 0
+        assert isinstance(max_buffer_size, int), (
+            "State should only ever be constructed with an integer or math.inf size."
+        )
+
+        self.max_buffer_size: float = max_buffer_size
+        self.buffer: mp.Queue[T | _MpEndOfStream] = mp.Queue(max_buffer_size)
+        self.closed: Event = mp.Event()
+
+    def __getstate__(self):
+        d = self.__dict__.copy()
+        d.pop("__orig_class__", None)
+        return d
+
+
+@dataclass(eq=False)
+class MpSender[T]:
+    """
+    An interprocess channel, mimicing the Anyio structure.
+    It should be noted that none of the clone methods are implemented for simplicity, for now.
+    """
+
+    _state: MpState[T] = field()
+
+    def send_nowait(self, item: T) -> None:
+        if self._state.closed.is_set():
+            raise ClosedResourceError
+        try:
+            self._state.buffer.put(item, block=False)
+        except Full:
+            raise WouldBlock from None
+        except ValueError as e:
+            print("Unreachable code path - let me know!")
+            raise ClosedResourceError from e
+
+    def send(self, item: T) -> None:
+        if self._state.closed.is_set():
+            raise ClosedResourceError
+        try:
+            self.send_nowait(item)
+        except WouldBlock:
+            # put anyway, blocking
+            self._state.buffer.put(item, block=True)
+
+    async def send_async(self, item: T) -> None:
+        await to_thread.run_sync(self.send, item, limiter=CapacityLimiter(1))
+
+    def close(self) -> None:
+        if not self._state.closed.is_set():
+            self._state.closed.set()
+        self._state.buffer.put(MP_END_OF_STREAM)
+        self._state.buffer.close()
+
+    # == context manager support ==#
+    def __enter__(self) -> Self:
+        return self
+
+    def __exit__(
+        self,
+        exc_type: type[BaseException] | None,
+        exc_val: BaseException | None,
+        exc_tb: TracebackType | None,
+    ) -> None:
+        self.close()
+
+    def __getstate__(self):
+        d = self.__dict__.copy()
+        d.pop("__orig_class__", None)
+        return d
+
+
+@dataclass(eq=False)
+class MpReceiver[T]:
+    """
+    An interprocess channel, mimicing the Anyio structure.
+    It should be noted that none of the clone methods are implemented for simplicity, for now.
+    """
+
+    _state: MpState[T] = field()
+
+    def receive_nowait(self) -> T:
+        if self._state.closed.is_set():
+            raise ClosedResourceError
+
+        try:
+            item = self._state.buffer.get(block=False)
+            if item is MP_END_OF_STREAM:
+                self.close()
+                raise EndOfStream
+            return item  # pyright: ignore[reportReturnType]
+        except Empty:
+            raise WouldBlock from None
+        except ValueError as e:
+            print("Unreachable code path - let me know!")
+            raise ClosedResourceError from e
+
+    def receive(self) -> T:
+        try:
+            return self.receive_nowait()
+        except WouldBlock:
+            item = self._state.buffer.get()
+            if item is MP_END_OF_STREAM:
+                self.close()
+                raise EndOfStream from None
+            return item  # pyright: ignore[reportReturnType]
+
+    async def receive_async(self) -> T:
+        return await to_thread.run_sync(self.receive, limiter=CapacityLimiter(1))
+
+    def close(self) -> None:
+        if not self._state.closed.is_set():
+            self._state.closed.set()
+        self._state.buffer.close()
+
+    # == iterator support ==#
+    def __iter__(self) -> Self:
+        return self
+
+    def __next__(self) -> T:
+        try:
+            return self.receive()
+        except EndOfStream:
+            raise StopIteration from None
+
+    # == async iterator support ==#
+    def __aiter__(self) -> Self:
+        return self
+
+    async def __anext__(self) -> T:
+        try:
+            return await self.receive_async()
+        except EndOfStream:
+            raise StopAsyncIteration from None
+
+    # == context manager support ==#
+    def __enter__(self) -> Self:
+        return self
+
+    def __exit__(
+        self,
+        exc_type: type[BaseException] | None,
+        exc_val: BaseException | None,
+        exc_tb: TracebackType | None,
+    ) -> None:
+        self.close()
+
+    def collect(self) -> list[T]:
+        """Collect all currently available items from this receiver"""
+        out: list[T] = []
+        while True:
+            try:
+                item = self.receive_nowait()
+                out.append(item)
+            except WouldBlock:
+                break
+        return out
+
+    def receive_at_least(self, n: int) -> list[T]:
+        out: list[T] = []
+        out.append(self.receive())
+        out.extend(self.collect())
+        while len(out) < n:
+            out.append(self.receive())
+            out.extend(self.collect())
+        return out
+
+    def __getstate__(self):
+        d = self.__dict__.copy()
+        d.pop("__orig_class__", None)
+        return d
+
+
 class channel[T]:  # noqa: N801
+    """Create a pair of asynchronous channels for communicating within the same process"""
+
     def __new__(cls, max_buffer_size: float = inf) -> tuple[Sender[T], Receiver[T]]:
         if max_buffer_size != inf and not isinstance(max_buffer_size, int):
             raise ValueError("max_buffer_size must be either an integer or math.inf")
         state = AnyioState[T](max_buffer_size)
         return Sender(_state=state), Receiver(_state=state)
+
+
+class mp_channel[T]:  # noqa: N801
+    """Create a pair of synchronous channels for interprocess communication"""
+
+    # max buffer size uses math.inf to represent an unbounded queue, and 0 to represent a yet unimplemented "unbuffered" queue.
+    def __new__(cls, max_buffer_size: float = inf) -> tuple[MpSender[T], MpReceiver[T]]:
+        if (
+            max_buffer_size == 0
+            or max_buffer_size != inf
+            and not isinstance(max_buffer_size, int)
+        ):
+            raise ValueError(
+                "max_buffer_size must be either an integer or math.inf. 0-sized buffers are not supported by multiprocessing"
+            )
+        state = MpState[T](max_buffer_size)
+        return MpSender(_state=state), MpReceiver(_state=state)
diff --git a/src/exo/utils/pydantic_ext.py b/src/exo/utils/pydantic_ext.py
index 5600d386..5631723c 100644
--- a/src/exo/utils/pydantic_ext.py
+++ b/src/exo/utils/pydantic_ext.py
@@ -37,3 +37,7 @@ class TaggedModel(CamelCaseModel):
             return handler(v[cls.__name__])
 
         return handler(v)
+
+    def __str__(self) -> str:
+        return f"{self.__class__.__name__}({super().__str__()})"
+
diff --git a/src/exo/utils/tests/testing_mp.py b/src/exo/utils/tests/testing_mp.py
new file mode 100644
index 00000000..62eddf0c
--- /dev/null
+++ b/src/exo/utils/tests/testing_mp.py
@@ -0,0 +1,41 @@
+import multiprocessing as mp
+import time
+
+import pytest
+from anyio import fail_after
+from loguru import logger
+
+from exo.utils.channels import MpReceiver, MpSender, mp_channel
+
+
+def foo(recv: MpReceiver[str]):
+    expected = ["hi", "hi 2", "bye"]
+    with recv as r:
+        for item in r:
+            assert item == expected.pop(0)
+
+
+def bar(send: MpSender[str]):
+    logger.warning("hi")
+    send.send("hi")
+    time.sleep(0.1)
+    logger.warning("hi 2")
+    send.send("hi 2")
+    time.sleep(0.1)
+    logger.warning("bye")
+    send.send("bye")
+    time.sleep(0.1)
+    send.close()
+
+
+# not async, just want the fail_after
+@pytest.mark.anyio
+async def test_channel_setup():
+    with fail_after(0.5):
+        s, r = mp_channel[str]()
+        p1 = mp.Process(target=foo, args=(r,))
+        p2 = mp.Process(target=bar, args=(s,))
+        p1.start()
+        p2.start()
+        p1.join()
+        p2.join()
diff --git a/src/exo/worker/common.py b/src/exo/worker/common.py
deleted file mode 100644
index 3f6517ba..00000000
--- a/src/exo/worker/common.py
+++ /dev/null
@@ -1,36 +0,0 @@
-from copy import deepcopy
-
-from pydantic import BaseModel, ConfigDict
-
-from exo.shared.types.common import Host
-from exo.shared.types.events import (
-    InstanceId,
-    RunnerStatusUpdated,
-)
-from exo.shared.types.worker.common import RunnerId
-from exo.shared.types.worker.runners import (
-    RunnerStatus,
-)
-from exo.shared.types.worker.shards import ShardMetadata
-from exo.worker.runner.runner_supervisor import RunnerSupervisor
-
-
-class AssignedRunner(BaseModel):
-    runner_id: RunnerId
-    instance_id: InstanceId
-    shard_metadata: ShardMetadata
-    hosts: list[Host] | None = None
-    mlx_ibv_devices: list[list[str | None]] | None = None
-    mlx_ibv_coordinator: str | None = None
-
-    status: RunnerStatus
-    failures: list[tuple[float, Exception]] = []
-    runner: RunnerSupervisor | None = None
-
-    model_config = ConfigDict(arbitrary_types_allowed=True)
-
-    def status_update_event(self) -> RunnerStatusUpdated:
-        return RunnerStatusUpdated(
-            runner_id=self.runner_id,
-            runner_status=deepcopy(self.status),
-        )
diff --git a/src/exo/worker/download/shard_downloader.py b/src/exo/worker/download/shard_downloader.py
index c5e557cb..a41b3eeb 100644
--- a/src/exo/worker/download/shard_downloader.py
+++ b/src/exo/worker/download/shard_downloader.py
@@ -26,7 +26,6 @@ class ShardDownloader(ABC):
 
         Args:
             shard (Shard): The shard to download.
-            inference_engine_name (str): The inference engine used on the node hosting the shard
         """
 
     @abstractmethod
diff --git a/src/exo/worker/main.py b/src/exo/worker/main.py
index 1091c807..31595c41 100644
--- a/src/exo/worker/main.py
+++ b/src/exo/worker/main.py
@@ -1,12 +1,7 @@
-import asyncio
-import time
-from asyncio import Queue
-from functools import partial
 from random import random
-from typing import AsyncGenerator
 
 import anyio
-from anyio import CancelScope, create_task_group
+from anyio import CancelScope, create_task_group, current_time
 from anyio.abc import TaskGroup
 from loguru import logger
 
@@ -15,53 +10,32 @@ from exo.shared.apply import apply
 from exo.shared.types.commands import ForwarderCommand, RequestEventLog
 from exo.shared.types.common import NodeId, SessionId
 from exo.shared.types.events import (
-    ChunkGenerated,
     Event,
     EventId,
     ForwarderEvent,
     IndexedEvent,
-    InstanceDeleted,
+    NodeDownloadProgress,
     NodeMemoryMeasured,
     NodePerformanceMeasured,
-    RunnerDeleted,
-    RunnerStatusUpdated,
-    TaskFailed,
-    TaskStateUpdated,
+    TaskCreated, TaskStatusUpdated,
     TopologyEdgeCreated,
     TopologyEdgeDeleted,
 )
 from exo.shared.types.multiaddr import Multiaddr
 from exo.shared.types.profiling import MemoryPerformanceProfile, NodePerformanceProfile
 from exo.shared.types.state import State
-from exo.shared.types.tasks import TaskId, TaskStatus
+from exo.shared.types.tasks import CreateRunner, DownloadModel, Task, TaskStatus, Shutdown
 from exo.shared.types.topology import Connection
-from exo.shared.types.worker.common import RunnerId
 from exo.shared.types.worker.downloads import (
     DownloadCompleted,
     DownloadOngoing,
     DownloadPending,
+    DownloadProgress,
 )
-from exo.shared.types.worker.ops import (
-    AssignRunnerOp,
-    ExecuteTaskOp,
-    RunnerDownOp,
-    RunnerFailedOp,
-    RunnerOp,
-    RunnerUpOp,
-    UnassignRunnerOp,
-)
-from exo.shared.types.worker.runners import (
-    DownloadingRunnerStatus,
-    FailedRunnerStatus,
-    InactiveRunnerStatus,
-    LoadedRunnerStatus,
-    RunningRunnerStatus,
-    StartingRunnerStatus,
-)
+from exo.shared.types.worker.runners import RunnerId
 from exo.shared.types.worker.shards import ShardMetadata
-from exo.utils.channels import Receiver, Sender
+from exo.utils.channels import Receiver, Sender, channel
 from exo.utils.event_buffer import OrderedBuffer
-from exo.worker.common import AssignedRunner
 from exo.worker.download.download_utils import (
     map_repo_download_progress_to_download_progress_data,
 )
@@ -80,12 +54,7 @@ class Worker:
         *,
         initial_connection_messages: list[ConnectionMessage],
         connection_message_receiver: Receiver[ConnectionMessage],
-        # Having written this pattern 3 times in the codebase:
-        # Should this be inherited??? Is this a real inheritance
-        # W????
-        # Limitation: This SHOULD be a MasterForwarderEvent, but inheritance says no :|
         global_event_receiver: Receiver[ForwarderEvent],
-        # Limitation: This SHOULD be a WorkerForwarderEvent, but inheritance says no :|
         local_event_sender: Sender[ForwarderEvent],
         # This is for requesting updates. It doesn't need to be a general command sender right now,
         # but I think it's the correct way to be thinking about commands
@@ -93,7 +62,10 @@ class Worker:
     ):
         self.node_id: NodeId = node_id
         self.session_id: SessionId = session_id
+
         self.shard_downloader: ShardDownloader = shard_downloader
+        self._pending_downloads: dict[RunnerId, ShardMetadata] = {}
+
         self.global_event_receiver = global_event_receiver
         self.local_event_sender = local_event_sender
         self.local_event_index = 0
@@ -104,10 +76,13 @@ class Worker:
         self.out_for_delivery: dict[EventId, ForwarderEvent] = {}
 
         self.state: State = State()
-        self.assigned_runners: dict[RunnerId, AssignedRunner] = {}
+        self.download_status: dict[ShardMetadata, DownloadProgress] = {}
+        self.runners: dict[RunnerId, RunnerSupervisor] = {}
         self._tg: TaskGroup | None = None
         self._nack_cancel_scope: CancelScope | None = None
 
+        self.event_sender, self.event_receiver = channel[Event]()
+
     async def run(self):
         logger.info("Starting Worker")
 
@@ -115,7 +90,7 @@ class Worker:
         async def resource_monitor_callback(
             node_performance_profile: NodePerformanceProfile,
         ) -> None:
-            await self.event_publisher(
+            await self.event_sender.send(
                 NodePerformanceMeasured(
                     node_id=self.node_id, node_profile=node_performance_profile
                 ),
@@ -124,7 +99,7 @@ class Worker:
         async def memory_monitor_callback(
             memory_profile: MemoryPerformanceProfile,
         ) -> None:
-            await self.event_publisher(
+            await self.event_sender.send(
                 NodeMemoryMeasured(node_id=self.node_id, memory=memory_profile)
             )
 
@@ -132,15 +107,17 @@ class Worker:
 
         async with create_task_group() as tg:
             self._tg = tg
+            tg.start_soon(self.plan_step)
             tg.start_soon(start_polling_node_metrics, resource_monitor_callback)
 
             tg.start_soon(start_polling_memory_metrics, memory_monitor_callback)
             tg.start_soon(self._connection_message_event_writer)
             tg.start_soon(self._resend_out_for_delivery)
             tg.start_soon(self._event_applier)
+            tg.start_soon(self._forward_events)
             # TODO: This is a little gross, but not too bad
             for msg in self._initial_connection_messages:
-                await self.event_publisher(
+                await self.event_sender.send(
                     self._convert_connection_message_to_event(msg)
                 )
             self._initial_connection_messages = []
@@ -148,9 +125,8 @@ class Worker:
         # Actual shutdown code - waits for all tasks to complete before executing.
         self.local_event_sender.close()
         self.command_sender.close()
-        for runner in self.assigned_runners.values():
-            if runner.runner:
-                await runner.runner.astop()
+        for runner in self.runners.values():
+            await runner.shutdown()
 
     async def _event_applier(self):
         with self.global_event_receiver as events:
@@ -162,14 +138,13 @@ class Worker:
 
                 # 2. for each event, apply it to the state
                 indexed_events = self.event_buffer.drain_indexed()
-                if not indexed_events:
-                    if (
-                        self._nack_cancel_scope is None
-                        or self._nack_cancel_scope.cancel_called
-                    ):
-                        assert self._tg
-                        self._tg.start_soon(self._nack_request)
-                elif self._nack_cancel_scope:
+                if not indexed_events and (
+                    self._nack_cancel_scope is None
+                    or self._nack_cancel_scope.cancel_called
+                ):
+                    assert self._tg
+                    self._tg.start_soon(self._nack_request)
+                elif indexed_events and self._nack_cancel_scope:
                     self._nack_cancel_scope.cancel()
 
                 flag = False
@@ -180,48 +155,82 @@ class Worker:
 
                 # 3. If we've found a "relevant" event, run a plan -> op -> execute cycle.
                 if flag:
-                    await self.plan_step()
+                    # await self.plan_step()
+                    pass
 
     async def plan_step(self):
-        # 3. based on the updated state, we plan & execute an operation.
-        op: RunnerOp | None = plan(
-            self.assigned_runners,
-            self.node_id,
-            self.state.instances,
-            self.state.runners,
-            self.state.tasks,
-        )
-
-        # run the op, synchronously blocking for now
-        if op is not None:
-            logger.info(f"Executing op {type(op)} {str(op)[:100]}")
-            logger.debug(f"Worker executing op: {type(op)} {str(op)[:100]}")
-            try:
-                async for event in self.execute_op(op):
-                    await self.event_publisher(event)
-            except Exception as e:
-                logger.opt(exception=e).warning(
-                    "Error occurred when executing task", flush=True
-                )
+        while True:
+            await anyio.sleep(0.1)
+            # 3. based on the updated state, we plan & execute an operation.
+            task: Task | None = plan(
+                self.node_id,
+                self.runners,
+                self.download_status,
+                self.state.downloads,
+                self.state.instances,
+                self.state.runners,
+                self.state.tasks,
+            )
+            if task is None:
+                continue
+            logger.info(f"Worker plan: {task.__class__.__name__}")
+            assert task.task_status
+            await self.event_sender.send(TaskCreated(task_id=task.task_id, task=task))
+
+            match task:
+                case CreateRunner():
+                    self._create_supervisor(task)
+                    await self.event_sender.send(TaskStatusUpdated(task_id=task.task_id, task_status=TaskStatus.Complete))
+                case DownloadModel(shard_metadata=shard):
+                    if shard not in self.download_status:
+                        progress = DownloadPending(
+                            shard_metadata=shard, node_id=self.node_id
+                        )
+                        self.download_status[shard] = progress
+                        await self.event_sender.send(
+                            NodeDownloadProgress(download_progress=progress)
+                        )
 
-                if isinstance(op, ExecuteTaskOp):
-                    generator = self.fail_task(
-                        e, runner_id=op.runner_id, task_id=op.task.task_id
+                    initial_progress = (
+                        await self.shard_downloader.get_shard_download_status_for_shard(
+                            shard
+                        )
                     )
-                else:
-                    generator = self.fail_runner(e, runner_id=op.runner_id)
-
-                async for event in generator:
-                    await self.event_publisher(event)
+                    if initial_progress.status == "complete":
+                        progress = DownloadCompleted(
+                            shard_metadata=shard, node_id=self.node_id
+                        )
+                        self.download_status[shard] = progress
+                        await self.event_sender.send(
+                            NodeDownloadProgress(download_progress=progress)
+                        )
+                        await self.event_sender.send(TaskStatusUpdated(task_id=task.task_id, task_status=TaskStatus.Complete))
+                    else:
+                        self.event_sender.send_nowait(TaskStatusUpdated(task_id=task.task_id, task_status=TaskStatus.Running))
+                        await self._handle_shard_download_process(
+                            task, initial_progress
+                        )
+                case Shutdown(runner_id=runner_id):
+                    await self.runners[runner_id].shutdown()
+                    del self.runners[runner_id]
+                case task:
+                    runner = self.runners[self._task_to_runner_id(task)]
+                    event = anyio.Event()
+                    await runner.start_task(task, event)
+                    await event.wait()
 
     def shutdown(self):
         if self._tg:
             self._tg.cancel_scope.cancel()
 
+    def _task_to_runner_id(self, task: Task):
+        instance = self.state.instances[task.instance_id]
+        return instance.shard_assignments.node_to_runner[self.node_id]
+
     async def _connection_message_event_writer(self):
         with self.connection_message_receiver as connection_messages:
             async for msg in connection_messages:
-                await self.event_publisher(
+                await self.event_sender.send(
                     self._convert_connection_message_to_event(msg)
                 )
 
@@ -278,377 +287,86 @@ class Worker:
 
     ## Op Executors
 
-    def _create_assigned_runner(self, op: AssignRunnerOp) -> AssignedRunner:
+    def _create_supervisor(self, task: CreateRunner) -> RunnerSupervisor:
         """Creates and stores a new AssignedRunner with initial downloading status."""
-        assigned_runner = AssignedRunner(
-            runner_id=op.runner_id,
-            instance_id=op.instance_id,
-            shard_metadata=op.shard_metadata,
-            hosts=op.hosts,
-            mlx_ibv_devices=op.mlx_ibv_devices,
-            mlx_ibv_coordinator=op.mlx_ibv_coordinator,
-            status=DownloadingRunnerStatus(
-                download_progress=DownloadPending(node_id=self.node_id)
-            ),
-            runner=None,
+        runner = RunnerSupervisor.create(
+            bound_instance=task.bound_instance,
+            event_sender=self.event_sender.clone(),
         )
-        self.assigned_runners[op.runner_id] = assigned_runner
-        return assigned_runner
-
-    async def _update_runner_status_to_completed_then_inactive(
-        self, assigned_runner: AssignedRunner
-    ) -> AsyncGenerator[Event, None]:
-        """Updates runner status from downloading to completed, then to inactive."""
-        assigned_runner.status = DownloadingRunnerStatus(
-            download_progress=DownloadCompleted(node_id=self.node_id)
-        )
-        yield assigned_runner.status_update_event()
-
-        assigned_runner.status = InactiveRunnerStatus()
-        yield assigned_runner.status_update_event()
-
-    async def _handle_already_downloaded_shard(
-        self, assigned_runner: AssignedRunner
-    ) -> AsyncGenerator[Event, None]:
-        """Handles the case where the shard is already downloaded."""
-        async for event in self._update_runner_status_to_completed_then_inactive(
-            assigned_runner
-        ):
-            yield event
+        self.runners[task.bound_instance.bound_runner_id] = runner
+        assert self._tg
+        self._tg.start_soon(runner.run)
+        return runner
 
     async def _handle_shard_download_process(
         self,
-        assigned_runner: AssignedRunner,
-        op: AssignRunnerOp,
+        task: DownloadModel,
         initial_progress: RepoDownloadProgress,
-    ) -> AsyncGenerator[Event, None]:
+    ):
         """Manages the shard download process with progress tracking."""
-        # Set initial ongoing status
-        assigned_runner.status = DownloadingRunnerStatus(
-            download_progress=DownloadOngoing(
-                node_id=self.node_id,
-                download_progress=map_repo_download_progress_to_download_progress_data(
-                    initial_progress
-                ),
-            )
+        status = DownloadOngoing(
+            node_id=self.node_id,
+            shard_metadata=task.shard_metadata,
+            download_progress=map_repo_download_progress_to_download_progress_data(
+                initial_progress
+            ),
         )
-        yield assigned_runner.status_update_event()
+        self.download_status[task.shard_metadata] = status
+        self.event_sender.send_nowait(NodeDownloadProgress(download_progress=status))
 
-        # Set up download progress tracking
-        download_progress_queue: asyncio.Queue[RepoDownloadProgress] = asyncio.Queue()
+        last_progress_time = 0.0
+        throttle_interval_secs = 1.0
 
+        # TODO: i hate callbacks
         def download_progress_callback(
             shard: ShardMetadata, progress: RepoDownloadProgress
         ) -> None:
-            download_progress_queue.put_nowait(progress)
-
-        self.shard_downloader.on_progress(download_progress_callback)
-        download_task = asyncio.create_task(
-            self.shard_downloader.ensure_shard(op.shard_metadata)
-        )
-
-        try:
-            async for event in self._monitor_download_progress(
-                assigned_runner, download_progress_queue
-            ):
-                yield event
-        finally:
-            if not download_task.done():
-                download_task.cancel()
-
-    async def _monitor_download_progress(
-        self,
-        assigned_runner: AssignedRunner,
-        download_progress_queue: asyncio.Queue[RepoDownloadProgress],
-    ) -> AsyncGenerator[Event, None]:
-        """Monitors download progress and yields status updates."""
-        last_progress_time = 0.0
-        throttle_interval_secs = 1.0
-
-        while True:
-            progress: RepoDownloadProgress = await asyncio.wait_for(
-                download_progress_queue.get(), timeout=15
-            )
-
+            nonlocal self
+            nonlocal last_progress_time
             if progress.status == "complete":
-                async for (
-                    event
-                ) in self._update_runner_status_to_completed_then_inactive(
-                    assigned_runner
-                ):
-                    yield event
-                break
-            elif progress.status == "in_progress":
-                if time.monotonic() - last_progress_time > throttle_interval_secs:
-                    assigned_runner.status = DownloadingRunnerStatus(
-                        download_progress=DownloadOngoing(
-                            node_id=self.node_id,
-                            download_progress=map_repo_download_progress_to_download_progress_data(
-                                progress
-                            ),
-                        )
-                    )
-                    yield assigned_runner.status_update_event()
-                    last_progress_time = time.monotonic()
-
-    async def _execute_assign_op(
-        self, op: AssignRunnerOp
-    ) -> AsyncGenerator[Event, None]:
-        """
-        A runner has been assigned. We need to also ensure that it's downloaded.
-        This op assigns the runner, and moves from Downloading -> Inactive (ready to spin) state.
-        """
-        assigned_runner = self._create_assigned_runner(op)
-        initial_progress = (
-            await self.shard_downloader.get_shard_download_status_for_shard(
-                op.shard_metadata
-            )
-        )
-
-        if initial_progress.status == "complete":
-            async for event in self._handle_already_downloaded_shard(assigned_runner):
-                yield event
-        else:
-            async for event in self._handle_shard_download_process(
-                assigned_runner, op, initial_progress
-            ):
-                yield event
-
-    async def _execute_unassign_op(
-        self, op: UnassignRunnerOp
-    ) -> AsyncGenerator[Event, None]:
-        if op.runner_id not in self.assigned_runners:
-            return
-
-        # We can try to do a graceful shutdown of the runner.
-        runner: RunnerSupervisor | None = self.assigned_runners[op.runner_id].runner
-        if runner is not None:
-            await runner.astop()
-
-        # This is all we really need:
-        del self.assigned_runners[op.runner_id]
-        yield RunnerDeleted(runner_id=op.runner_id)
-
-    async def _execute_runner_up_op(
-        self, op: RunnerUpOp, initialize_timeout: float | None = None
-    ) -> AsyncGenerator[Event, None]:
-        assigned_runner = self.assigned_runners[op.runner_id]
-
-        # Emit "Starting" status right away so UI can show loading state
-        assigned_runner.status = StartingRunnerStatus()
-        yield assigned_runner.status_update_event()
-
-        assigned_runner.runner = await RunnerSupervisor.create(
-            model_shard_meta=assigned_runner.shard_metadata,
-            hosts=assigned_runner.hosts,
-            mlx_ibv_devices=assigned_runner.mlx_ibv_devices,
-            mlx_ibv_coordinator=assigned_runner.mlx_ibv_coordinator,
-            initialize_timeout=initialize_timeout,
-        )
-
-        if assigned_runner.runner.runner_process.is_alive():
-            assigned_runner.status = LoadedRunnerStatus()
-        else:
-            runner = assigned_runner.runner
-            logger.warning(
-                f"Runner status is not runner_process.is_alive(): exit code {runner.runner_process.exitcode}"
-            )
-
-            assigned_runner.status = FailedRunnerStatus()
-        yield self.assigned_runners[op.runner_id].status_update_event()
-
-    async def _execute_runner_down_op(
-        self, op: RunnerDownOp
-    ) -> AsyncGenerator[Event, None]:
-        assigned_runner = self.assigned_runners[op.runner_id]
-
-        if isinstance(assigned_runner.runner, RunnerSupervisor):
-            await assigned_runner.runner.astop()
-
-        assigned_runner.runner = None
-
-        assigned_runner.status = InactiveRunnerStatus()
-        yield assigned_runner.status_update_event()
-        return
-
-    async def _execute_runner_failed_op(
-        self, op: RunnerFailedOp
-    ) -> AsyncGenerator[Event, None]:
-        """
-        We detected that this runner has failed. So we'll put it into 'failed' state now, triggering the rest of the instance to spin down.
-        """
-        assigned_runner = self.assigned_runners[op.runner_id]
-
-        if isinstance(assigned_runner.runner, RunnerSupervisor):
-            await (
-                assigned_runner.runner.astop()
-            )  # astop the runner to ensure it clears out of memory.
-
-        assigned_runner.status = FailedRunnerStatus()
-        yield self.assigned_runners[op.runner_id].status_update_event()
-
-    async def _execute_task_op(self, op: ExecuteTaskOp) -> AsyncGenerator[Event, None]:
-        """
-        This is the entry point for a chat completion starting.
-        While there is only one execute function, it will get called in different ways for runner 0 and runner [1, 2, 3, ...].
-        Runners [1, 2, 3, ...] will run this method when a task is in 'pending' state.
-        Runner 0 will run this method when a task is in 'running' state.
-        TODO: How do we handle the logic of ensuring that n-1 nodes have started their execution before allowing the 0'th runner to start?
-        This is still a little unclear to me.
-        """
-        assigned_runner = self.assigned_runners[op.runner_id]
-
-        async def inner_execute(queue: asyncio.Queue[Event]) -> None:
-            async def running_callback(queue: asyncio.Queue[Event]) -> None:
-                # Called when the MLX process has been kicked off
-                assigned_runner.status = RunningRunnerStatus()
-                await queue.put(assigned_runner.status_update_event())
-
-                if assigned_runner.shard_metadata.device_rank == 0:
-                    await queue.put(
-                        TaskStateUpdated(
-                            task_id=op.task.task_id,
-                            task_status=TaskStatus.Running,
-                        )
-                    )
-
-            assert assigned_runner.runner is not None
-            assert assigned_runner.runner.runner_process.is_alive()
-
-            async for chunk in assigned_runner.runner.stream_response(
-                task=op.task, request_started_callback=partial(running_callback, queue)
+                status = DownloadCompleted(shard_metadata=shard, node_id=self.node_id)
+                self.download_status[shard] = status
+                # Footgun!
+                self.event_sender.send_nowait(
+                    NodeDownloadProgress(download_progress=status)
+                )
+                self.event_sender.send_nowait(TaskStatusUpdated(task_id=task.task_id, task_status=TaskStatus.Complete))
+            elif (
+                progress.status == "in_progress"
+                and current_time() - last_progress_time > throttle_interval_secs
             ):
-                if assigned_runner.shard_metadata.device_rank == 0:
-                    await queue.put(
-                        ChunkGenerated(
-                            # TODO: at some point we will no longer have a bijection between task_id and row_id.
-                            # So we probably want to store a mapping between these two in our Worker object.
-                            command_id=chunk.command_id,
-                            chunk=chunk,
-                        )
-                    )
-
-            if op.task.task_id in self.state.tasks:
-                self.state.tasks[op.task.task_id].task_status = TaskStatus.Complete
-
-            if assigned_runner.shard_metadata.device_rank == 0:
-                # kind of hack - we don't want to wait for the round trip for this to complete
-                await queue.put(
-                    TaskStateUpdated(
-                        task_id=op.task.task_id,
-                        task_status=TaskStatus.Complete,
-                    )
+                status = DownloadOngoing(
+                    node_id=self.node_id,
+                    shard_metadata=shard,
+                    download_progress=map_repo_download_progress_to_download_progress_data(
+                        progress
+                    ),
                 )
-
-            # After a successful inference:
-            assigned_runner.status = LoadedRunnerStatus()
-            await queue.put(assigned_runner.status_update_event())
-
-        queue: Queue[Event] = asyncio.Queue()
-        task = asyncio.create_task(inner_execute(queue))
-
-        # TODO: Initial (prefil) timeout can be dynamic
-        # model_kb = assigned_runner.shard_metadata.model_meta.storage_size_kilobytes
-
-        try:
-            # Yield items from the queue
-            while True:
-                if task.done() and (exception := task.exception()):
-                    raise exception
-
-                try:
-                    # Use a timeout to periodically check task status
-                    item: Event = await asyncio.wait_for(queue.get(), timeout=0.01)
-                except asyncio.TimeoutError:
-                    continue
-
-                yield item
-                if isinstance(item, RunnerStatusUpdated) and isinstance(
-                    item.runner_status, (LoadedRunnerStatus, FailedRunnerStatus)
-                ):
-                    if isinstance(item.runner_status, LoadedRunnerStatus):
-                        assigned_runner.failures = []
-
-                    break
-        finally:
-            # Ensure the task is cleaned up
-            try:
-                await asyncio.wait_for(task, timeout=5)
-            except asyncio.TimeoutError:
-                logger.warning(
-                    "Timed out waiting for task cleanup after inference execution."
+                self.download_status[shard] = status
+                self.event_sender.send_nowait(
+                    NodeDownloadProgress(download_progress=status)
                 )
+                last_progress_time = current_time()
 
-    ## Operation Planner
-
-    async def execute_op(self, op: RunnerOp) -> AsyncGenerator[Event, None]:
-        ## It would be great if we can get rid of this async for ... yield pattern.
-        match op:
-            case AssignRunnerOp():
-                event_generator = self._execute_assign_op(op)
-            case UnassignRunnerOp():
-                event_generator = self._execute_unassign_op(op)
-            case RunnerUpOp():
-                event_generator = self._execute_runner_up_op(op)
-            case RunnerDownOp():
-                event_generator = self._execute_runner_down_op(op)
-            case RunnerFailedOp():
-                event_generator = self._execute_runner_failed_op(op)
-            case ExecuteTaskOp():
-                event_generator = self._execute_task_op(op)
-
-        async for event in event_generator:
-            yield event
-
-    async def fail_runner(
-        self, e: Exception, runner_id: RunnerId
-    ) -> AsyncGenerator[Event]:
-        if runner_id in self.assigned_runners:
-            assigned_runner = self.assigned_runners[runner_id]
-
-            if assigned_runner.runner is not None:
-                await assigned_runner.runner.astop()
-                assigned_runner.runner = None
-            assigned_runner.status = FailedRunnerStatus(error_message=str(e))
-            assigned_runner.failures.append((time.time(), e))
-
-            # Reset failure count back to 0 when succesful
-            if len(assigned_runner.failures) >= 3:
-                # Too many retries. We will emit a DeleteInstance
-                yield InstanceDeleted(instance_id=assigned_runner.instance_id)
-
-            yield assigned_runner.status_update_event()
-
-    async def fail_task(
-        self, e: Exception, runner_id: RunnerId, task_id: TaskId
-    ) -> AsyncGenerator[Event]:
-        if runner_id in self.assigned_runners:
-            yield TaskStateUpdated(
-                task_id=task_id,
-                task_status=TaskStatus.Failed,
-            )
-
-            yield TaskFailed(
-                task_id=task_id, error_type=str(type(e)), error_message=str(e)
-            )
-
-            async for event in self.fail_runner(e, runner_id):
-                yield event
+        self.shard_downloader.on_progress(download_progress_callback)
+        assert self._tg
+        self._tg.start_soon(self.shard_downloader.ensure_shard, task.shard_metadata)
 
-    # This function is re-entrant, take care!
-    async def event_publisher(self, event: Event) -> None:
-        fe = ForwarderEvent(
-            origin_idx=self.local_event_index,
-            origin=self.node_id,
-            session=self.session_id,
-            event=event,
-        )
-        logger.debug(
-            f"Worker published event {self.local_event_index}: {str(event)[:100]}"
-        )
-        self.local_event_index += 1
-        await self.local_event_sender.send(fe)
-        self.out_for_delivery[event.event_id] = fe
+    async def _forward_events(self) -> None:
+        with self.event_receiver as events:
+            async for event in events:
+                fe = ForwarderEvent(
+                    origin_idx=self.local_event_index,
+                    origin=self.node_id,
+                    session=self.session_id,
+                    event=event,
+                )
+                logger.debug(
+                    f"Worker published event {self.local_event_index}: {str(event)[:100]}"
+                )
+                self.local_event_index += 1
+                await self.local_event_sender.send(fe)
+                self.out_for_delivery[event.event_id] = fe
 
 
 def event_relevant_to_worker(event: Event, worker: Worker):
diff --git a/src/exo/worker/plan.py b/src/exo/worker/plan.py
index 8d0c7fa3..af46a3ff 100644
--- a/src/exo/worker/plan.py
+++ b/src/exo/worker/plan.py
@@ -1,303 +1,197 @@
-from typing import Mapping
+# pyright: reportUnusedImport = false
+
+from collections.abc import Mapping, Sequence
 
 from exo.shared.types.common import NodeId
-from exo.shared.types.events import (
-    InstanceId,
-)
-from exo.shared.types.tasks import Task, TaskId, TaskStatus
-from exo.shared.types.worker.common import RunnerId
-from exo.shared.types.worker.downloads import DownloadCompleted
-from exo.shared.types.worker.instances import Instance, InstanceStatus
-from exo.shared.types.worker.ops import (
-    AssignRunnerOp,
-    ExecuteTaskOp,
-    RunnerDownOp,
-    RunnerFailedOp,
-    RunnerOp,
-    RunnerUpOp,
-    UnassignRunnerOp,
+from exo.shared.types.tasks import (
+    ChatCompletion,
+    CreateRunner,
+    DownloadModel,
+    LoadModel,
+    Shutdown,
+    StartWarmup,
+    Task,
+    TaskId,
+    TaskStatus,
 )
+from exo.shared.types.worker.downloads import DownloadCompleted, DownloadProgress
+from exo.shared.types.worker.instances import BoundInstance, Instance, InstanceId
 from exo.shared.types.worker.runners import (
-    DownloadingRunnerStatus,
-    FailedRunnerStatus,
-    InactiveRunnerStatus,
-    LoadedRunnerStatus,
+    RunnerId,
+    RunnerLoaded,
+    RunnerLoading,
+    RunnerReady,
+    RunnerRunning,
     RunnerStatus,
-    RunningRunnerStatus,
-    StartingRunnerStatus,
+    RunnerWaitingForModel,
+    RunnerWarmingUp,
 )
-from exo.worker.common import AssignedRunner
+from exo.shared.types.worker.shards import ShardMetadata
+from exo.worker.runner.runner_supervisor import RunnerSupervisor
 
 
-def unassign_runners(
+def plan(
+    node_id: NodeId,
+    # Runners is expected to be FRESH and so should not come from state
+    runners: Mapping[RunnerId, RunnerSupervisor],
+    # DL_status is expected to be FRESH and so should not come from state
+    download_status: Mapping[ShardMetadata, DownloadProgress],
+    # gdls is not expected to be fresh
+    global_download_status: Mapping[NodeId, Sequence[DownloadProgress]],
     instances: Mapping[InstanceId, Instance],
-    state_runners: Mapping[RunnerId, RunnerStatus],
-    assigned_runners: dict[RunnerId, AssignedRunner],
-) -> UnassignRunnerOp | None:
-    runner_ids: set[RunnerId] = {
-        runner_id
-        for instance in instances.values()
-        for runner_id in instance.shard_assignments.runner_to_shard
-    }
-    for runner_id, _ in assigned_runners.items():
-        if runner_id not in runner_ids:
-            return UnassignRunnerOp(runner_id=runner_id)
+    all_runners: Mapping[RunnerId, RunnerStatus],  # all global
+    tasks: Mapping[TaskId, Task],
+) -> Task | None:
+    # Python short circuiting OR logic should evaluate these sequentially.
+    return (
+        _kill_runner(runners, all_runners, instances)
+        or _create_runner(node_id, runners, instances)
+        or _model_needs_download(runners, download_status)
+        or _load_model(runners, all_runners, global_download_status)
+        or _ready_to_warmup(runners, all_runners)
+        or _pending_tasks(runners, tasks, all_runners)
+    )
+
+
+def _kill_runner(
+    runners: Mapping[RunnerId, RunnerSupervisor],
+    all_runners: Mapping[RunnerId, RunnerStatus],
+    instances: Mapping[InstanceId, Instance],
+) -> Shutdown | None:
+    for runner in runners.values():
+        if (instance_id := runner.bound_instance.instance.instance_id) not in instances:
+            return Shutdown(instance_id=instance_id, runner_id = runner.bound_instance.bound_runner_id)
 
-    # If our instance is in 'downloading' or 'assigned' state, then we know the runner is stale. These are part of AssignRunnerOp and should be blocking.
-    for assigned_runner_id in assigned_runners:
-        if assigned_runner_id in state_runners:
-            status = state_runners[assigned_runner_id]
-            if isinstance(status, DownloadingRunnerStatus) and not isinstance(
-                status.download_progress, DownloadCompleted
-            ):
-                return UnassignRunnerOp(runner_id=assigned_runner_id)
+        """ --- Potential code to kill a runner if any runners in its instance have failed ---
+        global_runners_in_instance = runner.bound_instance.instance.shard_assignments.node_to_runner.values()
+        if any(isinstance(all_runners[runner_id], RunnerFailed) for runner_id in global_runners_in_instance if runner_id != runner.bound_instance.bound_runner_id):
+            Shutdown(instance_id=runner.bound_instance.instance.instance_id, runner_id=runner.bound_instance.bound_runner_id)
+        """
 
-    return None
 
+def _create_runner(
+    node_id: NodeId,
+    runners: Mapping[RunnerId, RunnerSupervisor],
+    instances: Mapping[InstanceId, Instance],
+) -> CreateRunner | None:
+    for instance in instances.values():
+        runner_id = instance.shard_assignments.node_to_runner.get(node_id, None)
+        if runner_id is None:
+            continue
 
-def failed_runners(
-    assigned_runners: dict[RunnerId, AssignedRunner],
-) -> RunnerFailedOp | None:
-    for runner_id, assigned_runner in assigned_runners.items():
-        if (
-            assigned_runner.runner is not None
-            and not assigned_runner.runner.runner_process.is_alive()
-            and not isinstance(assigned_runner.status, FailedRunnerStatus)
-        ):
-            return RunnerFailedOp(runner_id=runner_id)
-    return None
+        if runner_id in runners:
+            continue
 
+        shard = instance.shard(runner_id)
+        assert shard is not None
 
-def spin_down_runners(
-    instances: Mapping[InstanceId, Instance],
-    assigned_runners: dict[RunnerId, AssignedRunner],
-    state_runners: Mapping[RunnerId, RunnerStatus],
-    worker_node_id: NodeId,
-) -> RunnerDownOp | None:
-    for _instance_id, instance in instances.items():
-        for node_id, runner_id in instance.shard_assignments.node_to_runner.items():
-            if node_id != worker_node_id:
-                continue
+        return CreateRunner(
+            instance_id=instance.instance_id,
+            bound_instance=BoundInstance(instance=instance, bound_runner_id=runner_id),
+        )
 
-            # We spin down a runner if it's meant to be inactive and it's Loaded.
-            if (
-                runner_id in assigned_runners
-                and isinstance(assigned_runners[runner_id].status, LoadedRunnerStatus)
-                and instance.instance_type == InstanceStatus.Inactive
-            ):
-                return RunnerDownOp(runner_id=runner_id)
 
-    # If we are part of an instance that has a dead node - and we aren't the dead node - we should spin down
-    for _instance_id, instance in instances.items():
+def _model_needs_download(
+    runners: Mapping[RunnerId, RunnerSupervisor],
+    download_status: Mapping[ShardMetadata, DownloadProgress],
+) -> DownloadModel | None:
+    for runner in runners.values():
         if (
-            worker_node_id in instance.shard_assignments.node_to_runner
-            and instance.shard_assignments.node_to_runner[worker_node_id]
-            in assigned_runners
-            and not isinstance(
-                assigned_runners[
-                    instance.shard_assignments.node_to_runner[worker_node_id]
-                ].status,
-                InactiveRunnerStatus,
+            isinstance(runner.status, RunnerWaitingForModel)
+            and runner.bound_instance.bound_shard() not in download_status
+        ):
+            # We don't invalidate download_status randomly in case a file gets deleted on disk
+            return DownloadModel(
+                instance_id=runner.bound_instance.instance.instance_id,
+                shard_metadata=runner.bound_instance.bound_shard(),
             )
-        ):  # make sure that our runner has not already been spun down into ready state
-            other_node_in_instance_has_failed = False
-            for runner_id in instance.shard_assignments.runner_to_shard:
-                if (
-                    runner_id in state_runners
-                    and isinstance(state_runners[runner_id], FailedRunnerStatus)
-                    and runner_id not in assigned_runners
-                ):
-                    other_node_in_instance_has_failed = True
 
-            if other_node_in_instance_has_failed:
-                # Spin down *our* runner
-                return RunnerDownOp(
-                    runner_id=instance.shard_assignments.node_to_runner[worker_node_id]
-                )
 
-    # If we are failed - and *all of the other nodes have spun down* - then we can spin down too.
-    for _instance_id, instance in instances.items():
+""" --- TODO!
+def _init_backend(
+    runners: Mapping[RunnerId, RunnerSupervisor],
+    all_runners: Mapping[RunnerId, RunnerStatus],
+) -> LoadModel | None:
+    for runner in runner.values()
+    pass
+"""
+
+
+def _load_model(
+    runners: Mapping[RunnerId, RunnerSupervisor],
+    all_runners: Mapping[RunnerId, RunnerStatus],
+    global_download_status: Mapping[NodeId, Sequence[DownloadProgress]],
+) -> LoadModel | None:
+    for runner in runners.values():
         if (
-            worker_node_id in instance.shard_assignments.node_to_runner
-            and instance.shard_assignments.node_to_runner[worker_node_id]
-            in state_runners
-            and instance.shard_assignments.node_to_runner[worker_node_id]
-            in assigned_runners
-            and isinstance(
-                assigned_runners[
-                    instance.shard_assignments.node_to_runner[worker_node_id]
-                ].status,
-                FailedRunnerStatus,
+            all(
+                isinstance(dp, DownloadCompleted)
+                if dp.shard_metadata
+                == runner.bound_instance.instance.shard_assignments.runner_to_shard[rid]
+                else True
+                for nid, rid in runner.bound_instance.instance.shard_assignments.node_to_runner.items()
+                for dp in global_download_status[nid]
+            )
+            and isinstance(runner.status, RunnerWaitingForModel)
+            and all(
+                isinstance(
+                    all_runners.get(global_runner_id, None),
+                    (RunnerWaitingForModel, RunnerLoading, RunnerLoaded),
+                )
+                for global_runner_id in runner.bound_instance.instance.shard_assignments.runner_to_shard
             )
         ):
-            num_spundown_nodes = 0
-            for runner_id in instance.shard_assignments.runner_to_shard:
-                if (
-                    runner_id in state_runners
-                    and isinstance(state_runners[runner_id], InactiveRunnerStatus)
-                    and runner_id not in assigned_runners
-                ):
-                    num_spundown_nodes += 1
-                # Suggested:
-                # if runner_id in state_runners and isinstance(state.runners[runner_id], InactiveRunnerStatus):
-                #     if runner_id != instance.shard_assignments.node_to_runner[worker_node_id]:
-                #         num_spundown_nodes += 1
-
-            if (
-                num_spundown_nodes
-                == next(
-                    iter(instance.shard_assignments.runner_to_shard.values())
-                ).world_size
-                - 1
-            ):
-                # All the other nodes are spun down - so now we can spin down too.
-                # This also catches the case of 1-node. If there's one node in the instance then we should spin down straight away
-                return RunnerDownOp(
-                    runner_id=instance.shard_assignments.node_to_runner[worker_node_id]
+            return LoadModel(instance_id=runner.bound_instance.instance.instance_id)
+
+
+def _ready_to_warmup(
+    runners: Mapping[RunnerId, RunnerSupervisor],
+    all_runners: Mapping[RunnerId, RunnerStatus],
+) -> StartWarmup | None:
+    for runner in runners.values():
+        if isinstance(runner.status, RunnerLoaded) and (
+            (
+                all(
+                    isinstance(
+                        all_runners.get(global_runner_id, None),
+                        (RunnerLoaded, RunnerWarmingUp),
+                    )
+                    for global_runner_id in runner.bound_instance.instance.shard_assignments.runner_to_shard
                 )
-    return None
-
-
-def assign_runners(
-    instances: Mapping[InstanceId, Instance],
-    assigned_runners: dict[RunnerId, AssignedRunner],
-    worker_node_id: NodeId,
-) -> AssignRunnerOp | None:
-    for instance_id, instance in instances.items():
-        for node_id, runner_id in instance.shard_assignments.node_to_runner.items():
-            if node_id != worker_node_id:
-                continue
-
-            if runner_id not in assigned_runners:
-                return AssignRunnerOp(
-                    runner_id=runner_id,
-                    instance_id=instance_id,
-                    shard_metadata=instance.shard_assignments.runner_to_shard[
-                        runner_id
-                    ],
-                    hosts=instance.hosts,
-                    mlx_ibv_devices=instance.mlx_ibv_devices,
-                    mlx_ibv_coordinator=instance.mlx_ibv_coordinator,
+                and runner.bound_instance.bound_shard().device_rank != 0
+            )
+            or (
+                all(
+                    isinstance(
+                        all_runners.get(global_runner_id, None), (RunnerWarmingUp)
+                    )
+                    for global_runner_id in runner.bound_instance.instance.shard_assignments.runner_to_shard
+                    if global_runner_id != runner.bound_instance.bound_runner_id
                 )
-    return None
-
-
-def spin_up_runners(
-    instances: Mapping[InstanceId, Instance],
-    assigned_runners: dict[RunnerId, AssignedRunner],
-    state_runners: Mapping[RunnerId, RunnerStatus],
-    worker_node_id: NodeId,
-) -> RunnerUpOp | None:
-    for _instance_id, instance in instances.items():
-        if (
-            worker_node_id in instance.shard_assignments.node_to_runner
-            and assigned_runners[
-                instance.shard_assignments.node_to_runner[worker_node_id]
-            ].runner
-            is None
-            and instance.instance_type == InstanceStatus.Active
+                and runner.bound_instance.bound_shard().device_rank == 0
+            )
         ):
-            # We are part of this instance, we want it up but it hasn't been spun up yet.
-            # Need to assert all other runners are ready before we can spin up.
-            ready_to_spin = True
-            for runner_id in instance.shard_assignments.node_to_runner.values():
-                if runner_id in state_runners and isinstance(
-                    state_runners[runner_id],
-                    (
-                        InactiveRunnerStatus,
-                        StartingRunnerStatus,
-                    ),
-                ):
-                    ready_to_spin = False
-
-            if ready_to_spin:
-                return RunnerUpOp(
-                    runner_id=instance.shard_assignments.node_to_runner[worker_node_id]
-                )
-    return None
+            return StartWarmup(instance_id=runner.bound_instance.instance.instance_id)
 
 
-def execute_task_op(
-    instances: Mapping[InstanceId, Instance],
-    assigned_runners: dict[RunnerId, AssignedRunner],
-    state_runners: Mapping[RunnerId, RunnerStatus],
+def _pending_tasks(
+    runners: Mapping[RunnerId, RunnerSupervisor],
     tasks: Mapping[TaskId, Task],
-    worker_node_id: NodeId,
-) -> ExecuteTaskOp | None:
-    for instance_id, instance in instances.items():
-        for node_id, runner_id in instance.shard_assignments.node_to_runner.items():
-            if node_id != worker_node_id:
+    all_runners: Mapping[RunnerId, RunnerStatus],
+) -> Task | None:
+    for task in tasks.values():
+        # for now, just forward chat completions
+        if not isinstance(task, ChatCompletion):
+            continue
+        if task.task_status not in (TaskStatus.Pending, TaskStatus.Running):
+            continue
+
+        for runner in runners.values():
+            if task.instance_id != runner.bound_instance.instance.instance_id:
                 continue
-            assert runner_id in assigned_runners
-            runner = assigned_runners[runner_id]
-            if not isinstance(runner.status, LoadedRunnerStatus):
-                continue  # The only previous state to get to Running is from Loaded
-
-            for _, task in tasks.items():
-                if task.instance_id == instance_id and (
-                    task.task_status in (TaskStatus.Pending, TaskStatus.Failed)
-                ):
-                    if (
-                        runner.shard_metadata.device_rank >= 1
-                        or runner.shard_metadata.world_size == 1
-                    ):
-                        return ExecuteTaskOp(runner_id=runner_id, task=task)
-                    else:
-                        # We already know our own status is Loaded. We are rank 0,
-                        # so let's check that all the other runners are running - ready for us to fire the prompt.
-                        running_runner_count = 0
-                        for (
-                            other_runner_id,
-                            other_runner_status,
-                        ) in state_runners.items():
-                            if (
-                                other_runner_id
-                                in instance.shard_assignments.node_to_runner.values()
-                                and isinstance(other_runner_status, RunningRunnerStatus)
-                            ):
-                                running_runner_count += 1
-
-                        if running_runner_count == runner.shard_metadata.world_size - 1:
-                            return ExecuteTaskOp(runner_id=runner_id, task=task)
 
-    return None
-
-
-def plan(
-    assigned_runners: dict[RunnerId, AssignedRunner],
-    worker_node_id: NodeId,
-    instances: Mapping[InstanceId, Instance],
-    state_runners: Mapping[RunnerId, RunnerStatus],  # all global
-    tasks: Mapping[TaskId, Task],
-) -> RunnerOp | None:
-    # First, unassign assigned runners that are no longer in the state.
-    if unop := unassign_runners(instances, state_runners, assigned_runners):
-        return unop
-
-    # mark failed runners that are not marked yet as failed
-    if failed_op := failed_runners(assigned_runners):
-        return failed_op
-
-    # spin down runners that are no longer needed
-    if down_op := spin_down_runners(
-        instances, assigned_runners, state_runners, worker_node_id
-    ):
-        return down_op
-
-    # Then assign runners we do want
-    if assign_op := assign_runners(instances, assigned_runners, worker_node_id):
-        return assign_op
-
-    # Then spin up 'ready' runners that should be active
-    if runner_up_op := spin_up_runners(
-        instances, assigned_runners, state_runners, worker_node_id
-    ):
-        return runner_up_op
-
-    # Then make sure things are running based on tasks.
-    if exec_op := execute_task_op(
-        instances, assigned_runners, state_runners, tasks, worker_node_id
-    ):
-        return exec_op
-
-    return None
+            if isinstance(runner.status, RunnerReady) and all(
+                isinstance(all_runners[global_runner_id], (RunnerReady, RunnerRunning))
+                for global_runner_id in runner.bound_instance.instance.shard_assignments.runner_to_shard
+            ):
+                return task
diff --git a/src/exo/worker/runner/bootstrap.py b/src/exo/worker/runner/bootstrap.py
index bc734155..989b8723 100644
--- a/src/exo/worker/runner/bootstrap.py
+++ b/src/exo/worker/runner/bootstrap.py
@@ -1,10 +1,17 @@
-import asyncio
+"""--- not doing this anymore
 import faulthandler
 import os
 import sys
-from multiprocessing.connection import Connection
+"""
 
+import loguru
 
+from exo.shared.types.events import Event
+from exo.shared.types.tasks import Task
+from exo.shared.types.worker.instances import BoundInstance
+from exo.utils.channels import MpReceiver, MpSender
+
+""" -- not doing this anymore
 def _redirect_stderr_to_file(path: str) -> None:
     # Replace fd 2 (stderr) with a file descriptor pointing to `path`
     fd = os.open(path, os.O_WRONLY | os.O_CREAT | os.O_TRUNC, 0o644)
@@ -12,20 +19,36 @@ def _redirect_stderr_to_file(path: str) -> None:
     os.close(fd)
     # Rebind sys.stderr so Python's own writes go to the new fd as well (line-buffered)
     sys.stderr = os.fdopen(2, "w", buffering=1, closefd=False)
+"""
 
 
-def entrypoint(raw_conn: Connection, err_path: str) -> None:
+def entrypoint(
+    bound_instance: BoundInstance,
+    event_sender: MpSender[Event],
+    task_receiver: MpReceiver[Task],
+    # err_path: str,
+    _logger: "loguru.Logger",
+) -> None:
     """
     Minimal entrypoint for the spawned child process.
 
     It redirects fd=2 (stderr) to a pipe provided by the parent, *then* imports
     the heavy runner module so that any C/C++ or MLX logs/crashes land in that pipe.
     """
-    # os.environ["MLX_METAL_FAST_SYNCH"] = "1"
+    """ --- not doing this anymore
     _redirect_stderr_to_file(err_path)
     faulthandler.enable(file=sys.stderr, all_threads=True)
+    """
+    import os
+    os.environ["MLX_METAL_FAST_SYNCH"] = "1"
+
+    global logger
+    logger = _logger
 
     # Import the heavy runner only after stderr is redirected
     from exo.worker.runner.runner import main
 
-    asyncio.run(main(raw_conn))
+    main(bound_instance, event_sender, task_receiver)
+
+
+logger: "loguru.Logger"
diff --git a/src/exo/worker/runner/generate.py b/src/exo/worker/runner/generate.py
index 4f91e9e8..1293184c 100644
--- a/src/exo/worker/runner/generate.py
+++ b/src/exo/worker/runner/generate.py
@@ -1,35 +1,24 @@
-import asyncio
-import concurrent.futures
-import functools
-import time
-from collections.abc import AsyncGenerator
-from functools import partial
-from typing import Any, Callable, Generator
+from typing import Any, Callable, Generator, get_args, cast
 
 import mlx.core as mx
-from mlx.core import array
-from mlx_lm.models import cache
 from mlx_lm.models.cache import KVCache
+from mlx_lm import stream_generate
 
-from exo.engines.mlx import Model, TokenizerWrapper
+from exo.engines.mlx import Model
+from mlx_lm.tokenizer_utils import TokenizerWrapper
 from exo.engines.mlx.utils_mlx import (
-    apply_chat_template,
-    broadcast_from_zero,
     make_kv_cache,
+    apply_chat_template,
     mx_barrier,
 )
+from exo.shared.openai_compat import FinishReason
 from exo.shared.types.api import ChatCompletionMessage
 from exo.shared.types.tasks import ChatCompletionTaskParams
 from exo.shared.types.worker.commands_runner import (
     GenerationResponse,
-    RunnerMessage,
-    RunnerResponse,
     TokenizedResponse,
 )
-from exo.shared.types.worker.communication import (
-    AsyncConnection,
-    runner_print,
-)
+from exo.worker.runner.bootstrap import logger
 
 generation_stream = mx.new_stream(mx.default_device())
 
@@ -48,281 +37,15 @@ def maybe_quantize_kv_cache(
         ):
             prompt_cache[e] = c.to_quantized(group_size=kv_group_size, bits=kv_bits)
 
-
-def generate_step(
-    prompt: mx.array,
-    model: Model,
-    *,
-    max_tokens: int = 256,
-    sampler: Callable[[mx.array], mx.array],
-    logits_processors: list[Callable[[mx.array, mx.array], mx.array]] | None = None,
-    max_kv_size: int | None = None,
-    prompt_cache: list[KVCache] | None = None,
-    prefill_step_size: int = 16384,
-    kv_bits: int | None = None,
-    kv_group_size: int = 64,
-    quantized_kv_start: int = 0,
-    prompt_progress_callback: Callable[[int, int], None] | None = None,
-    input_embeddings: mx.array | None = None,
-    group: mx.distributed.Group | None = None,
-) -> Generator[tuple[int, mx.array], None, None]:
-    """
-    A generator producing token ids based on the given prompt from the model.
-
-    Args:
-        prompt (mx.array): The input prompt.
-        model (Model): The model to use for generation.
-        max_tokens (int): The maximum number of tokens. Use``-1`` for an infinite
-          generator. Default: ``256``.
-        sampler (Callable[mx.array, mx.array]): A sampler for sampling a
-          token from a vector of log probabilities.
-        logits_processors (list[Callable[[mx.array, mx.array], mx.array]], optional):
-          A list of functions that take tokens and logits and return the processed
-          logits. Default: ``None``.
-        max_kv_size (int, optional): Maximum size of the key-value cache. Old
-          entries (except the first 4 tokens) will be overwritten.
-        prompt_cache (list[Any], optional): A pre-computed prompt cache. Note, if
-          provided, the cache will be updated in place.
-        prefill_step_size (int): Step size for processing the prompt.
-        kv_bits (int, optional): Number of bits to use for KV cache quantization.
-          None implies no cache quantization. Default: ``None``.
-        kv_group_size (int): Group size for KV cache quantization. Default: ``64``.
-        quantized_kv_start (int): Step to begin using a quantized KV cache.
-           when ``kv_bits`` is non-None. Default: ``0``.
-        prompt_progress_callback (Callable[[int, int], None]): A call-back which takes the
-           prompt tokens processed so far and the total number of prompt tokens.
-        input_embeddings (mx.array, optional): Input embeddings to use instead of or in
-          conjunction with prompt tokens. Default: ``None``.
-
-    Yields:
-        tuple[int, mx.array]: One token and a vector of log probabilities.
-    """
-    if input_embeddings is not None:
-        if len(prompt) > 0 and len(prompt) != len(input_embeddings):
-            raise ValueError(
-                f"When providing input_embeddings, their sequence length ({len(input_embeddings)}) "
-                f"must match the sequence length of the prompt ({len(prompt)}), or the "
-                "prompt must be empty."
-            )
-    elif len(prompt) == 0:
-        raise ValueError(
-            "Either input_embeddings or prompt (or both) must be provided."
-        )
-
-    tokens = None
-
-    if prompt_cache is None:
-        prompt_cache = cache.make_prompt_cache(
-            model,
-            max_kv_size=max_kv_size,
-        )
-
-    prompt_progress_callback = prompt_progress_callback or (lambda _, __: None)
-
-    quantize_cache_fn = functools.partial(
-        maybe_quantize_kv_cache,
-        quantized_kv_start=quantized_kv_start,
-        kv_group_size=kv_group_size,
-        kv_bits=kv_bits,
-    )
-
-    def _model_call(
-        input_tokens: mx.array, input_embeddings: mx.array | None
-    ) -> mx.array:
-        if input_embeddings is not None:
-            return model(
-                input_tokens,
-                cache=prompt_cache,
-                input_embeddings=input_embeddings,
-            )
-        else:
-            return model(input_tokens, cache=prompt_cache)
-
-    def _step(
-        input_tokens: mx.array, input_embeddings: mx.array | None = None
-    ) -> tuple[mx.array, mx.array]:
-        nonlocal tokens
-
-        with mx.stream(generation_stream):
-            logits = _model_call(
-                input_tokens=input_tokens[None],
-                input_embeddings=(
-                    input_embeddings[None] if input_embeddings is not None else None
-                ),
-            )
-
-            logits = logits[:, -1, :]
-
-            if logits_processors and len(input_tokens) > 0:
-                tokens = (
-                    mx.concat([tokens, input_tokens])
-                    if tokens is not None
-                    else input_tokens
-                )
-                for processor in logits_processors:
-                    logits = processor(tokens, logits)
-
-            quantize_cache_fn(prompt_cache)
-
-            logprobs = logits - mx.logsumexp(logits, keepdims=True)
-            sampled = sampler(logprobs)
-            return sampled, logprobs.squeeze(0)
-
-    with mx.stream(generation_stream):
-        total_prompt_tokens = (
-            len(input_embeddings) if input_embeddings is not None else len(prompt)
-        )
-        prompt_processed_tokens = 0
-        prompt_progress_callback(prompt_processed_tokens, total_prompt_tokens)
-
-        while total_prompt_tokens - prompt_processed_tokens > prefill_step_size:
-            runner_print(
-                f"Prefilling {min(prefill_step_size, len(prompt))} tokens. Remaining tokens: {len(prompt)}. Peak memory: {mx.get_peak_memory() // 2**30} GB"
-            )
-            n_to_process = min(prefill_step_size, prompt.size)
-            _model_call(
-                input_tokens=prompt[:n_to_process][None],
-                input_embeddings=(
-                    input_embeddings[:n_to_process][None]
-                    if input_embeddings is not None
-                    else None
-                ),
-            )
-            quantize_cache_fn(prompt_cache)
-
-            start_time = time.time()
-            mx.eval([c.state for c in prompt_cache])
-            eval_time = time.time() - start_time
-            prompt_processed_tokens += n_to_process
-
-            prompt = prompt[n_to_process:]
-            input_embeddings = (
-                input_embeddings[n_to_process:]
-                if input_embeddings is not None
-                else input_embeddings
-            )
-
-            mx.clear_cache()
-            # if eval_time > 7.0:
-            #     prefill_step_size = prefill_step_size // 2
-            if group is not None:
-                prefill_step_size = broadcast_from_zero(prefill_step_size)
-            prefill_step_size = max(1, prefill_step_size)
-            prompt_progress_callback(prompt_processed_tokens, total_prompt_tokens)
-
-        if prompt_processed_tokens > 0:
-            runner_print("finished prefill stage.")
-
-        y, logprobs = _step(input_tokens=prompt, input_embeddings=input_embeddings)
-
-    mx.async_eval(y, logprobs)
-    next_y: array | None = None
-    next_logprobs: array | None = None
-    n = 0
-
-    while True:
-        assert y is not None
-        assert logprobs is not None
-        if n != max_tokens:
-            next_y, next_logprobs = _step(y)
-            mx.async_eval(next_y, next_logprobs)
-        if n == 0:
-            mx.eval(y)
-            prompt_progress_callback(total_prompt_tokens, total_prompt_tokens)
-        if n == max_tokens:
-            break
-        yield int(y.item()), logprobs
-        if n % 256 == 0:
-            mx.clear_cache()
-        y, logprobs = next_y, next_logprobs
-        n += 1
-
-
-def stream_generate(
+def warmup_inference(
     model: Model,
     tokenizer: TokenizerWrapper,
-    prompt: str,
-    max_tokens: int,
     sampler: Callable[[mx.array], mx.array],
-    conn: AsyncConnection[RunnerResponse, RunnerMessage] | None,
-    logits_processors: list[Callable[[mx.array, mx.array], mx.array]] | None = None,
-    max_kv_size: int | None = None,
-    prompt_cache: list[KVCache] | None = None,
-    prefill_step_size: int = 2048,
-    kv_bits: int | None = None,
-    kv_group_size: int = 64,
-    quantized_kv_start: int = 0,
-    prompt_progress_callback: Callable[[int, int], None] | None = None,
-    input_embeddings: mx.array | None = None,
-    group: mx.distributed.Group | None = None,
-) -> Generator[GenerationResponse, None, None]:
-    # Try to infer if special tokens are needed
-    add_special_tokens = tokenizer.bos_token is None or not prompt.startswith(
-        tokenizer.bos_token
-    )
-    prompt_array: mx.array = mx.array(
-        tokenizer.encode(prompt, add_special_tokens=add_special_tokens)
-    )
-    if conn is not None:
-        conn.send_sync(TokenizedResponse(prompt_tokens=len(prompt_array)))
-
-    detokenizer = tokenizer.detokenizer
-
-    token_generator: Generator[tuple[int, array], None, None] = generate_step(
-        prompt_array,
-        model,
-        max_tokens=max_tokens,
-        sampler=sampler,
-        logits_processors=logits_processors,
-        max_kv_size=max_kv_size,
-        prompt_cache=prompt_cache,
-        prefill_step_size=prefill_step_size,
-        kv_bits=kv_bits,
-        kv_group_size=kv_group_size,
-        quantized_kv_start=quantized_kv_start,
-        prompt_progress_callback=prompt_progress_callback,
-        input_embeddings=input_embeddings,
-        group=group,
-    )
-
-    token = None
-    detokenizer.reset()
-    for token, _ in token_generator:
-        if token in tokenizer.eos_token_ids:
-            break
-
-        detokenizer.add_token(token)
-
-        # TODO: We could put more metrics on this GenerationResponse if we wish
-        yield GenerationResponse(
-            text=detokenizer.last_segment,
-            token=token,
-            finish_reason=None,
-        )
-
-    assert token is not None
-    detokenizer.finalize()
-    yield GenerationResponse(
-        text=detokenizer.last_segment,
-        token=token,
-        finish_reason="stop" if token in tokenizer.eos_token_ids else "length",
-    )
-
-
-async def warmup_inference(
-    mlx_executor: concurrent.futures.ThreadPoolExecutor,
-    model: Model,
-    tokenizer: TokenizerWrapper,
-    sampler: Callable[[mx.array], mx.array],
-    group: mx.distributed.Group | None = None,
 ) -> int:
-    loop = asyncio.get_running_loop()
-
-    warmup_prompt = await apply_chat_template(
-        mlx_executor=mlx_executor,
+    warmup_prompt = apply_chat_template(
         tokenizer=tokenizer,
         chat_task_data=ChatCompletionTaskParams(
-            model="warmup",
+            model="",
             messages=[
                 ChatCompletionMessage(
                     role="user",
@@ -334,95 +57,63 @@ async def warmup_inference(
 
     tokens_generated = 0
 
-    def _generate_warmup():
-        nonlocal tokens_generated
-        runner_print("Generating warmup tokens")
-        for _r in stream_generate(
-            model=model,
-            tokenizer=tokenizer,
-            prompt=warmup_prompt,
-            max_tokens=50,
-            sampler=sampler,
-            conn=None,
-            group=group,
-        ):
-            runner_print("Generated warmup token: " + str(_r.text))
-            tokens_generated += 1
+    cache = make_kv_cache(
+        model=model,
+    )
+
+    logger.info("Generating warmup tokens")
+    for _r in stream_generate(
+        model=model,
+        tokenizer=tokenizer,
+        prompt=warmup_prompt,
+        max_tokens=50,
+        sampler=sampler,
+        prompt_cache=cache,
+        prefill_step_size=65536,
+    ):
+        logger.info("Generated warmup token: " + str(_r.text))
+        tokens_generated += 1
 
-    await loop.run_in_executor(mlx_executor, _generate_warmup)
-    runner_print("Generated ALL warmup tokens")
-    await loop.run_in_executor(mlx_executor, lambda: mx_barrier(group))
+    logger.info("Generated ALL warmup tokens")
+    mx_barrier()
 
     return tokens_generated
 
 
-async def mlx_generate(
-    mlx_executor: concurrent.futures.ThreadPoolExecutor,
+def mlx_generate(
     model: Model,
     tokenizer: TokenizerWrapper,
     sampler: Callable[[mx.array], mx.array],
     task: ChatCompletionTaskParams,
-    conn: AsyncConnection[RunnerResponse, RunnerMessage],
-) -> AsyncGenerator[GenerationResponse]:
-    loop = asyncio.get_running_loop()
-    queue: asyncio.Queue[GenerationResponse | Exception | object] = asyncio.Queue()
-    sentinel = object()
-
-    def _generate_tokens(prompt: str, max_tokens: int, cache: list[KVCache]) -> None:
-        try:
-            for generation_response in stream_generate(
-                model=model,
-                tokenizer=tokenizer,
-                prompt=prompt,
-                max_tokens=max_tokens,
-                sampler=sampler,
-                prompt_cache=cache,
-                prefill_step_size=1024,
-                conn=conn,
-            ):
-                _ = loop.call_soon_threadsafe(queue.put_nowait, generation_response)
-        except Exception as e:
-            _ = loop.call_soon_threadsafe(queue.put_nowait, e)
-        finally:
-            _ = loop.call_soon_threadsafe(queue.put_nowait, sentinel)
-
+) -> Generator[GenerationResponse]:
     # Currently we support chat-completion tasks only.
-    runner_print(f"task_params: {task}")
+    logger.info(f"task_params: {task}")
 
-    prompt = await apply_chat_template(
-        mlx_executor=mlx_executor,
+    prompt = apply_chat_template(
         tokenizer=tokenizer,
         chat_task_data=task,
     )
 
-    cache_future = loop.run_in_executor(
-        mlx_executor,
-        lambda: asyncio.run(
-            make_kv_cache(
-                model=model,
-            )
-        ),
+    cache = make_kv_cache(
+        model=model,
     )
-    cache = await cache_future
 
     max_tokens = task.max_tokens or 1000
-    generation_fn = partial(_generate_tokens, prompt, max_tokens, cache)
-
-    future = loop.run_in_executor(mlx_executor, generation_fn)
-
-    while True:
-        item = await queue.get()
-        queue.task_done()
-
-        if item is sentinel:
-            break
-
-        if isinstance(item, Exception):
-            raise item
-
-        assert isinstance(item, GenerationResponse)  # constrain datatype
-        runner_print(item.text)
-        yield item
+    for out in stream_generate(
+        model=model,
+        tokenizer=tokenizer,
+        prompt=prompt,
+        max_tokens=max_tokens,
+        sampler=sampler,
+        prompt_cache=cache,
+        prefill_step_size=65536,
+    ):
+        logger.info(out.text)
+        if out.finish_reason != None and out.finish_reason not in get_args(FinishReason):
+            # We don't throw here as this failure case is really not all that bad
+            # Just log the error and move on
+            logger.warning(f"Model generated unexpected finish_reason: {out.finish_reason}")
 
-    # Wait for the executor thread to complete
-    await future
+        yield GenerationResponse(
+            text=out.text, token=out.token, finish_reason=cast(FinishReason | None, out.finish_reason)
+        )
diff --git a/src/exo/worker/runner/runner.py b/src/exo/worker/runner/runner.py
index 79b9b521..0cbbe079 100644
--- a/src/exo/worker/runner/runner.py
+++ b/src/exo/worker/runner/runner.py
@@ -1,127 +1,218 @@
-import asyncio
-import concurrent.futures
 import time
-from functools import partial
-from multiprocessing.connection import Connection
 
 from exo.engines.mlx.utils_mlx import (
     mx_barrier,
     initialize_mlx,
     mlx_force_oom,
 )
-from exo.shared.global_conn import set_conn
+from exo.shared.types.chunks import TokenChunk
+from exo.shared.types.events import (
+    ChunkGenerated,
+    Event,
+    RunnerStatusUpdated,
+    TaskAcknowledged,
+    TaskStatusUpdated,
+)
+from exo.shared.types.tasks import (
+    ChatCompletion,
+    LoadModel,
+    Shutdown,
+    StartWarmup,
+    Task,
+    TaskStatus,
+)
 from exo.shared.types.worker.commands_runner import (
-    ChatTaskMessage,
-    ExitMessage,
-    FinishedResponse,
-    InitializedResponse,
-    RunnerMessage,
-    RunnerResponse,
-    SetupMessage,
+    GenerationResponse,
+    TokenizedResponse,
 )
-from exo.shared.types.worker.communication import (
-    AsyncConnection,
-    runner_print,
-    runner_write_error,
+from exo.shared.types.worker.instances import BoundInstance
+from exo.shared.types.worker.runners import (
+    RunnerFailed,
+    RunnerLoaded,
+    RunnerLoading,
+    RunnerReady,
+    RunnerRunning,
+    RunnerStatus,
+    RunnerWaitingForModel,
+    RunnerWarmingUp,
 )
-from exo.shared.types.worker.shards import ShardMetadata
-from exo.utils import ensure_type
-from exo.worker.runner.generate import mlx_generate
-
+from exo.utils.channels import MpReceiver, MpSender
+from exo.worker.runner.bootstrap import logger
+from exo.worker.runner.generate import mlx_generate, warmup_inference
 
-async def main(raw_conn: Connection):
-    conn = AsyncConnection[RunnerResponse, RunnerMessage](raw_conn)
-    set_conn(conn)
 
+def main(
+    bound_instance: BoundInstance,
+    event_sender: MpSender[Event],
+    task_receiver: MpReceiver[Task],
+):
+    instance, runner_id, shard_metadata = (
+        bound_instance.instance,
+        bound_instance.bound_runner_id,
+        bound_instance.bound_shard(),
+    )
     try:
-        runner_print("hello from the runner")
-        init_message = await conn.recv()
-        setup_message = ensure_type(init_message, SetupMessage)
-        model_shard_meta: ShardMetadata = setup_message.model_shard_meta
-        hosts = setup_message.hosts
-        mlx_ibv_devices = setup_message.mlx_ibv_devices
-        mlx_ibv_coordinator = setup_message.mlx_ibv_coordinator
-
-        if getattr(model_shard_meta, "immediate_exception", False):
+        logger.info("hello from the runner")
+        if getattr(shard_metadata, "immediate_exception", False):
             raise Exception("Fake exception - runner failed to spin up.")
-        if timeout := getattr(model_shard_meta, "should_timeout", 0):
-            await asyncio.sleep(timeout)
+        if timeout := getattr(shard_metadata, "should_timeout", 0):
+            time.sleep(timeout)
 
         setup_start_time = time.time()
 
-        mlx_executor = concurrent.futures.ThreadPoolExecutor(max_workers=1)
-        loop = asyncio.get_running_loop()
+        model = None
+        tokenizer = None
+        sampler = None
 
-        model, tokenizer, sampler, group = await loop.run_in_executor(
-            mlx_executor,
-            partial(
-                initialize_mlx,
-                model_shard_meta=model_shard_meta,
-                hosts=hosts,
-                mlx_ibv_devices=mlx_ibv_devices,
-                mlx_ibv_coordinator=mlx_ibv_coordinator,
-            ),
+        current_status: RunnerStatus = RunnerWaitingForModel()
+        logger.info("runner waiting for model")
+        event_sender.send(
+            RunnerStatusUpdated(runner_id=runner_id, runner_status=current_status)
         )
+        with task_receiver as tasks:
+            for task in tasks:
+                event_sender.send(
+                    TaskStatusUpdated(
+                        task_id=task.task_id, task_status=TaskStatus.Running
+                    )
+                )
+                event_sender.send(TaskAcknowledged(task_id=task.task_id))
+                match task:
+                    case LoadModel() if isinstance(
+                        current_status, (RunnerWaitingForModel, RunnerFailed)
+                    ):
+                        current_status = RunnerLoading()
+                        logger.info("runner loading")
+                        event_sender.send(
+                            RunnerStatusUpdated(
+                                runner_id=runner_id, runner_status=current_status
+                            )
+                        )
 
-        # runner_print(
-        #     f"Warming up inference for model_shard_meta: {model_shard_meta} hosts: {hosts}"
-        # )
-        # toks = await warmup_inference(
-        #     mlx_executor=mlx_executor,
-        #     model=model,
-        #     tokenizer=tokenizer,
-        #     sampler=sampler,
-        #     group=group,
-        # )
-        # runner_print(f"Warmed up by generating {toks} tokens")
-        runner_print("Synchronizing processes before generation")
-        await loop.run_in_executor(mlx_executor, lambda: mx_barrier(group))
-        runner_print("Synchronized processes before generation")
-        await conn.send(InitializedResponse(time_taken=time.time() - setup_start_time))
+                        model, tokenizer, sampler = initialize_mlx(bound_instance)
 
-        while True:
-            message = await conn.recv()
-            match message:
-                case ChatTaskMessage(task_data=task):
-                    runner_print(f"received chat request: {str(task)[:500]}")
-                    # Ensure we have a chat-completion task subtype
-                    # TODO: this is a hack, why are we only looking at the first message? should have a tokenizer
-                    prompt = task.messages[0]
-                    if (
-                        prompt.content is not None
-                        and "EXO RUNNER MUST FAIL" in prompt.content
-                    ):
-                        runner_print("raising exception")
-                        raise Exception(
-                            "Artificial runner exception - for testing purposes only."
+                        current_status = RunnerLoaded()
+                        logger.info("runner loaded")
+                        event_sender.send(
+                            RunnerStatusUpdated(
+                                runner_id=runner_id, runner_status=current_status
+                            )
+                        )
+                    case StartWarmup() if isinstance(current_status, RunnerLoaded):
+                        assert model
+                        assert tokenizer
+                        assert sampler
+                        current_status = RunnerWarmingUp()
+                        logger.info("runner warming up")
+                        event_sender.send(
+                            RunnerStatusUpdated(
+                                runner_id=runner_id, runner_status=current_status
+                            )
                         )
-                    if (
-                        prompt.content is not None
-                        and "EXO RUNNER MUST OOM" in prompt.content
-                    ):
-                        mlx_force_oom()
-                    if (
-                        prompt.content is not None
-                        and "EXO RUNNER MUST TIMEOUT" in prompt.content
-                    ):
-                        await asyncio.sleep(100)
 
-                    # Generate responses using the actual MLX generation
-                    async for generation_response in mlx_generate(
-                        mlx_executor=mlx_executor,
-                        model=model,
-                        tokenizer=tokenizer,
-                        sampler=sampler,
-                        task=task,
-                        conn=conn,
-                    ):
-                        await conn.send(generation_response)
+                        logger.info(f"warming up inference for instance: {instance}")
+                        toks = warmup_inference(
+                            model=model,
+                            tokenizer=tokenizer,
+                            sampler=sampler,
+                        )
+                        logger.info(f"warmed up by generating {toks} tokens")
+                        logger.info(
+                            f"runner initialized in {time.time() - setup_start_time} seconds"
+                        )
+                        current_status = RunnerReady()
+                        logger.info("runner ready")
+                        event_sender.send(
+                            RunnerStatusUpdated(
+                                runner_id=runner_id, runner_status=RunnerReady()
+                            )
+                        )
+                    case ChatCompletion(
+                        task_params=task_params, command_id=command_id
+                    ) if isinstance(current_status, RunnerReady):
+                        assert model
+                        assert tokenizer
+                        assert sampler
+                        logger.info(f"received chat request: {str(task)[:500]}")
+                        current_status = RunnerRunning()
+                        logger.info("runner running")
+                        event_sender.send(
+                            RunnerStatusUpdated(
+                                runner_id=runner_id, runner_status=current_status
+                            )
+                        )
+                        # Ensure we have a chat-completion task subtype
+                        # TODO: this is a hack, why are we only looking at the first message? should have a tokenizer
+                        prompt = task_params.messages[0]
+                        if (
+                            prompt.content is not None
+                            and "EXO RUNNER MUST FAIL" in prompt.content
+                        ):
+                            logger.info("raising exception")
+                            raise Exception(
+                                "Artificial runner exception - for testing purposes only."
+                            )
+                        if (
+                            prompt.content is not None
+                            and "EXO RUNNER MUST OOM" in prompt.content
+                        ):
+                            mlx_force_oom()
+                        if (
+                            prompt.content is not None
+                            and "EXO RUNNER MUST TIMEOUT" in prompt.content
+                        ):
+                            time.sleep(100)
 
-                    await conn.send(FinishedResponse())
-                case ExitMessage():
-                    break
-                case _:
-                    raise ValueError(f"Unknown message: {message}")
+                        # Generate responses using the actual MLX generation
+                        for response in mlx_generate(
+                            model=model,
+                            tokenizer=tokenizer,
+                            sampler=sampler,
+                            task=task_params,
+                        ):
+                            match response:
+                                case GenerationResponse():
+                                    if shard_metadata.device_rank == 0:
+                                        event_sender.send(
+                                            ChunkGenerated(
+                                                command_id=command_id,
+                                                chunk=TokenChunk(
+                                                    idx=response.token,
+                                                    model=shard_metadata.model_meta.model_id,
+                                                    text=response.text,
+                                                    token_id=response.token,
+                                                    finish_reason=response.finish_reason,
+                                                ),
+                                            )
+                                        )
+                                case TokenizedResponse():
+                                    # TODO: something here ig
+                                    logger.info("Finished tokenizing?")
+
+                        current_status = RunnerReady()
+                        logger.info("runner ready")
+                        event_sender.send(
+                            RunnerStatusUpdated(
+                                runner_id=runner_id, runner_status=RunnerReady()
+                            )
+                        )
+                    case Shutdown():
+                        break
+                    case _:
+                        raise ValueError("Received task outside of state machine")
+                event_sender.send(
+                    TaskStatusUpdated(
+                        task_id=task.task_id, task_status=TaskStatus.Complete
+                    )
+                )
 
     except Exception as e:
-        runner_write_error(e)
+        logger.opt(exception=e).warning(
+            f"Runner {runner_id} crashed with critical exception {e}"
+        )
+        event_sender.send(
+            RunnerStatusUpdated(
+                runner_id=runner_id,
+                runner_status=RunnerFailed(error_message=str(e)),
+            )
+        )
diff --git a/src/exo/worker/runner/runner_supervisor.py b/src/exo/worker/runner/runner_supervisor.py
index 7a7fe8a9..785f9f11 100644
--- a/src/exo/worker/runner/runner_supervisor.py
+++ b/src/exo/worker/runner/runner_supervisor.py
@@ -1,334 +1,184 @@
-import asyncio
 import contextlib
-import multiprocessing as mp
-import os
 import signal
-import tempfile
-import traceback
+import sys
+from dataclasses import dataclass, field
 from multiprocessing import Process
-from multiprocessing.connection import Connection
-from typing import Any, AsyncGenerator, Callable, Coroutine
+from typing import Self
 
+import anyio
 import psutil
+from anyio import (
+    BrokenResourceError,
+    ClosedResourceError,
+    EndOfStream,
+    create_task_group,
+    current_time,
+    to_thread,
+)
+from anyio.abc import TaskGroup
 from loguru import logger
 
-from exo.shared.global_conn import (
-    AsyncConnection,
-)
-from exo.shared.types.chunks import GenerationChunk, TokenChunk
-from exo.shared.types.common import CommandId, Host
-from exo.shared.types.tasks import ChatCompletionTaskParams, Task
-from exo.shared.types.worker.commands_runner import (
-    ChatTaskMessage,
-    ErrorResponse,
-    FinishedResponse,
-    GenerationResponse,
-    InitializedResponse,
-    PrintResponse,
-    RunnerMessage,
-    RunnerResponse,
-    SetupMessage,
-    TokenizedResponse,
+from exo.shared.types.events import Event, RunnerStatusUpdated, TaskAcknowledged
+from exo.shared.types.tasks import Task, TaskId
+from exo.shared.types.worker.instances import BoundInstance
+from exo.shared.types.worker.runners import (
+    RunnerError,
+    RunnerFailed,
+    RunnerStatus,
+    RunnerWaitingForModel,
 )
-from exo.shared.types.worker.common import RunnerError
 from exo.shared.types.worker.shards import ShardMetadata
+from exo.utils.channels import MpReceiver, MpSender, Sender, mp_channel
 from exo.worker.runner.bootstrap import entrypoint
 from exo.worker.runner.utils import (
     get_weights_size,
 )
 
-INITIALIZE_TIMEOUT = 400
 PREFILL_TIMEOUT_SECONDS = 60
 DECODE_TIMEOUT_SECONDS = 5
 
 
+@dataclass(eq=False)
 class RunnerSupervisor:
-    def __init__(
-        self,
-        model_shard_meta: ShardMetadata,
-        hosts: list[Host] | None,
-        mlx_ibv_devices: list[list[str | None]] | None,
-        mlx_ibv_coordinator: str | None,
-        runner_process: Process,
-        conn: Connection,
-        read_queue: asyncio.Queue[RunnerResponse],
-        err_path: str,
-    ):
-        self.model_shard_meta = model_shard_meta
-        self.hosts = hosts
-        self.mlx_ibv_devices = mlx_ibv_devices
-        self.mlx_ibv_coordinator = mlx_ibv_coordinator
-        self.runner_process = runner_process
-
-        self.conn = AsyncConnection[RunnerMessage, RunnerResponse](conn)
-        self._raw_conn = conn
-
-        self.read_queue = read_queue
-        self.read_task = asyncio.create_task(self._read_coro())
-
-        self.err_path = err_path
+    shard_metadata: ShardMetadata
+    bound_instance: BoundInstance
+    runner_process: Process
+    initialize_timeout: float
+    _ev_recv: MpReceiver[Event]
+    _task_sender: MpSender[Task]
+    _event_sender: Sender[Event]
+    # err_path: str
+    _tg: TaskGroup | None = field(default=None, init=False)
+    status: RunnerStatus = field(default_factory=RunnerWaitingForModel, init=False)
+    pending: dict[TaskId, anyio.Event] = field(default_factory=dict, init=False)
 
     @classmethod
-    async def create(
+    def create(
         cls,
-        model_shard_meta: ShardMetadata,
-        hosts: list[Host] | None = None,
-        mlx_ibv_devices: list[list[str | None]] | None = None,
-        mlx_ibv_coordinator: str | None = None,
-        initialize_timeout: float | None = None,
-    ) -> "RunnerSupervisor":
-        """
-        Create and initialize a RunnerSupervisor instance.
-        The .create() classmethod pattern is used to ensure the constructor is asynchronous.
-        """
-        ctx = mp.get_context("spawn")
-        parent_conn, child_conn = ctx.Pipe(duplex=True)
-
+        *,
+        bound_instance: BoundInstance,
+        event_sender: Sender[Event],
+        initialize_timeout: float = 400,
+    ) -> Self:
+        ev_send, ev_recv = mp_channel[Event]()
+        # A task is kind of a runner command
+        task_sender, task_recv = mp_channel[Task]()
+
+        """ --- not doing this for now
         with tempfile.NamedTemporaryFile(
             prefix="child_stderr_", suffix=".log", delete=False
         ) as tmp:
             err_path = tmp.name
 
+        """
         runner_process = Process(
-            target=entrypoint, args=(child_conn, err_path), daemon=False
+            target=entrypoint,
+            args=(
+                bound_instance,
+                ev_send,
+                task_recv,
+                # err_path,
+                logger,
+            ),
+            daemon=True,
         )
-        runner_process.start()
-        child_conn.close()
 
-        read_queue = asyncio.Queue[RunnerResponse]()
+        shard_metadata = bound_instance.bound_shard()
 
         self = cls(
-            model_shard_meta=model_shard_meta,
-            hosts=hosts,
-            mlx_ibv_devices=mlx_ibv_devices,
-            mlx_ibv_coordinator=mlx_ibv_coordinator,
+            bound_instance=bound_instance,
+            shard_metadata=shard_metadata,
             runner_process=runner_process,
-            read_queue=read_queue,
-            conn=parent_conn,
-            err_path=err_path,
+            initialize_timeout=initialize_timeout,
+            _ev_recv=ev_recv,
+            _task_sender=task_sender,
+            _event_sender=event_sender,
+            # err_path=err_path,
         )
 
-        logger.info(f"Initializing mlx instance with {model_shard_meta=}")
-        await self.conn.send(
-            SetupMessage(
-                model_shard_meta=model_shard_meta,
-                hosts=hosts,
-                mlx_ibv_devices=mlx_ibv_devices,
-                mlx_ibv_coordinator=mlx_ibv_coordinator,
-            )
-        )
-
-        initialize_timeout = initialize_timeout or INITIALIZE_TIMEOUT
-        response = await self._read_with_error_check(timeout=initialize_timeout)
-
-        assert isinstance(response, InitializedResponse)
-        logger.info(f"Runner initialized in {response.time_taken} seconds")
-
         return self
 
-    async def _read_with_error_check(self, timeout: float) -> RunnerResponse | None:
-        """
-        Read from the queue with a timeout, but also check if the read_task has failed.
-        """
-        if self.read_task.done():
-            e = self.read_task.exception()
-            await self.astop()
-            if e is not None:
-                raise e
-            else:
-                return None
-
-        queue_task = asyncio.create_task(self.read_queue.get())
-
-        done, pending = await asyncio.wait(
-            [queue_task, self.read_task],
-            timeout=timeout,
-            return_when=asyncio.FIRST_COMPLETED,
-        )
-
-        for task in pending:
-            if task is queue_task:
-                task.cancel()
-
-        if queue_task in done:
-            return await queue_task
-
-        if self.read_task in done:
-            await self.astop()
-            await self.read_task  # Re-raises any exception from read_task
-
-            # This should never get hit.
-            raise RunnerError(
-                "RunnerStopped",
-                "Runner read loop terminated unexpectedly before any response.",
-                "",
-            )
-
-        # if we haven't read from the queue, we have timed out.
-        await self.astop()  # TODO: This could be handled by the called or _read_with_error_check - as we don't want a false Timeout to bring the whole runner down.
-        raise asyncio.TimeoutError()
-
-    async def _read_coro(self):
-        while True:
-            try:
-                response: RunnerResponse = await self.conn.recv()
-            except EOFError as e_eof:
-                e = await self._raise_crashed()
-                if e is not None:
-                    raise e from e_eof
-                break
-
-            match response:
-                case PrintResponse():
-                    # TODO: THIS IS A REALLY IMPORTANT LOG MESSAGE, AND SHOULD BE MADE PRETTIER
-                    logger.info(f"{response.text}")
-                case ErrorResponse():
-                    raise RunnerError(
-                        response.error_type, response.error_message, response.traceback
-                    )
-                case _:
-                    await self.read_queue.put(response)
-
-    async def stream_response(
-        self,
-        task: Task,
-        request_started_callback: Callable[..., Coroutine[Any, Any, None]]
-        | None = None,
-    ) -> AsyncGenerator[GenerationChunk, None]:
-        """
-        Streams a chat request from the model.
-        The request is pushed to the runner, and if the shard is the terminal shard, the response is streamed back to the worker.
-        request_started_callback is called once the request is pushed to the runner, used to publish InferencePrepareCompleted and InferenceTriggerCompleted events.
-        """
-        if not self.runner_process.is_alive():
-            raise RuntimeError("Runner process was found to be dead")
-
-        task_params = task.task_params
-        assert isinstance(
-            task_params, ChatCompletionTaskParams
-        )  # this is messy for now.
-        await self.conn.send(
-            ChatTaskMessage(
-                task_data=task_params,
-            ),
-        )
-
-        response = await self._read_with_error_check(5.0)
-        assert isinstance(response, TokenizedResponse)
-
-        if request_started_callback is not None:
-            await request_started_callback()
-
-        timeout = PREFILL_TIMEOUT_SECONDS
-
-        logger.info(f"Starting chat completion with timeout {timeout}")
-
-        while True:
-            try:
-                response = await self._read_with_error_check(timeout)
-            except asyncio.TimeoutError as e:
-                logger.error(
-                    f"Generation timed out during {'prefill' if timeout == PREFILL_TIMEOUT_SECONDS else 'decoding stage'}"
-                )
-                raise e
-
-            match response:
-                case GenerationResponse():
-                    yield TokenChunk(
-                        command_id=CommandId(task.command_id),
-                        idx=response.token,
-                        model=self.model_shard_meta.model_meta.model_id,
-                        text=response.text,
-                        token_id=response.token,
-                        finish_reason=response.finish_reason,
-                    )
-                    timeout = DECODE_TIMEOUT_SECONDS
-                case FinishedResponse():
-                    break
-                case _:
-                    raise ValueError(f"Unexpected response type found: {response}")
-
-    async def astop(self) -> None:
-        # Cancel the stderr monitoring task
-        async def await_task(task: asyncio.Task[Any]):
-            if not task.done():
-                task.cancel()
-                with contextlib.suppress(asyncio.CancelledError):
-                    await task
-
-        await await_task(self.read_task)
+    async def run(self):
+        self.runner_process.start()
+        async with create_task_group() as tg:
+            self._tg = tg
+            tg.start_soon(self._forward_events)
 
+        self._ev_recv.close()
+        self._task_sender.close()
+        self._event_sender.close()
         self.runner_process.kill()
+        await to_thread.run_sync(self.runner_process.join)
 
-        with contextlib.suppress(Exception):
-            self._raw_conn.close()
+    async def start_task(self, task: Task, event: anyio.Event):
+        self.pending[task.task_id] = event
+        self._task_sender.send(task)
 
-        # Wait to make sure that the model has been unloaded from memory
-        async def wait_for_memory_release() -> None:
-            required_memory_bytes = get_weights_size(self.model_shard_meta).in_bytes
-            start_time = asyncio.get_event_loop().time()
+    async def _forward_events(self):
+        with self._ev_recv as events:
             while True:
-                available_memory_bytes = psutil.virtual_memory().available
-                if available_memory_bytes >= required_memory_bytes:
-                    break
-                if asyncio.get_event_loop().time() - start_time > 30.0:
-                    logger.warning(
-                        "Runner memory not released after 30 seconds - exiting"
-                    )
+                try:
+                    event = await events.receive_async()
+                except (ClosedResourceError, BrokenResourceError, EndOfStream):
+                    await self._check_runner()
                     break
-                await asyncio.sleep(0.1)
-
-        await wait_for_memory_release()
+                if isinstance(event, RunnerStatusUpdated):
+                    self.status = event.runner_status
+                if isinstance(event, TaskAcknowledged):
+                    self.pending.pop(event.task_id).set()
+                    continue
+                await self._event_sender.send(event)
+
+    async def shutdown(self) -> None:
+        assert self._tg
+        self._tg.cancel_scope.cancel()
+
+        required_memory_bytes = get_weights_size(self.shard_metadata).in_bytes
+        start_time = current_time()
+        while True:
+            available_memory_bytes = psutil.virtual_memory().available
+            if available_memory_bytes >= required_memory_bytes:
+                break
+            if current_time() - start_time > 30.0:
+                logger.warning("Runner memory not released after 30 seconds - exiting")
+                break
+            await anyio.sleep(1)
 
     def __del__(self) -> None:
         if self.runner_process.is_alive():
-            logger.warning(
-                "RunnerSupervisor was not stopped cleanly before garbage collection. Force killing process tree."
-            )
-            # Can't use async in __del__, so use psutil directly
-            try:
-                pid = self.runner_process.pid
-                if pid:
-                    parent = psutil.Process(pid)
-                    children = parent.children(recursive=True)
-                    for child in reversed(children):
-                        with contextlib.suppress(
-                            psutil.NoSuchProcess, psutil.AccessDenied
-                        ):
-                            child.kill()
-                    with contextlib.suppress(psutil.NoSuchProcess, psutil.AccessDenied):
-                        parent.kill()
-            except Exception:
-                with contextlib.suppress(ProcessLookupError):
-                    self.runner_process.kill()
-
-    async def _raise_crashed(self) -> Exception | None:
-        await asyncio.sleep(0.1)
+            logger.warning("RunnerSupervisor was not stopped cleanly.")
+            with contextlib.suppress(ValueError):
+                self.runner_process.kill()
 
+    async def _check_runner(self) -> RunnerError | None:
         rc = self.runner_process.exitcode
         if rc == 0:
-            return None
+            logger.warning("Runner closed communication without terminating process")
 
+        """ --- not doing this anymore
         try:
             with open(self.err_path, "r", errors="replace") as f:
                 captured = f.read()
         finally:
             with contextlib.suppress(OSError):
                 os.unlink(self.err_path)
+        """
 
-        # 2) Describe cause (signal vs exitcode)
-        cause = f"exitcode={rc}"
         if isinstance(rc, int) and rc < 0:
             sig = -rc
             try:
                 cause = f"signal={sig} ({signal.strsignal(sig)})"
             except Exception:
                 cause = f"signal={sig}"
+        else:
+            cause = f"exitcode={rc}"
 
-        logger.error(f"Runner terminated ({cause}).\n{captured}")
+        logger.opt(exception=sys.exception()).error(f"Runner terminated ({cause})")
 
-        return RunnerError(
-            error_type="RunnerCrash",
-            error_message=f"Runner terminated ({cause}).\n{captured}",
-            traceback=traceback.format_exc(),
+        await self._event_sender.send(
+            RunnerStatusUpdated(
+                runner_id=self.bound_instance.bound_runner_id,
+                runner_status=RunnerFailed(error_message=f"Terminated ({cause})"),
+            )
         )
+        await self.shutdown()
diff --git a/src/exo/worker/runner/utils.py b/src/exo/worker/runner/utils.py
index 1242390d..9cf22c95 100644
--- a/src/exo/worker/runner/utils.py
+++ b/src/exo/worker/runner/utils.py
@@ -6,7 +6,7 @@ import psutil
 from loguru import logger
 
 from exo.shared.types.memory import Memory
-from exo.shared.types.worker.shards import ShardMetadata
+from exo.shared.types.worker.shards import PipelineShardMetadata, ShardMetadata
 
 
 async def kill_process_tree(runner_process: asyncio.subprocess.Process) -> None:
@@ -58,7 +58,7 @@ def get_weights_size(model_shard_meta: ShardMetadata) -> Memory:
         * model_shard_meta.model_meta.storage_size.in_kb
         / (
             1
-            if model_shard_meta.strategy in ["auto", "pipeline", "pipeline_rdma"]
+            if isinstance(model_shard_meta, PipelineShardMetadata)
             else model_shard_meta.world_size
         )
     )
diff --git a/tmp/run_llm.sh b/tmp/run_llm.sh
index b08db159..07599c2d 100755
--- a/tmp/run_llm.sh
+++ b/tmp/run_llm.sh
@@ -13,7 +13,7 @@ QUERY="$*"
 curl -sN -X POST "http://$HOST:8000/v1/chat/completions" \
   -H "Content-Type: application/json" \
   -d "{
-        \"model\": \"mlx-community/DeepSeek-V3.1-8bit\",
+        \"model\": \"mlx-community/Llama-3.3-70B-Instruct-8bit\",
         \"stream\": true,
         \"messages\": [{ \"role\": \"user\",   \"content\": \"$QUERY\" }]
       }" |
@@ -21,4 +21,4 @@ curl -sN -X POST "http://$HOST:8000/v1/chat/completions" \
   grep --line-buffered -v 'data: \[DONE\]' |
   cut -d' ' -f2- |
   jq -r --unbuffered '.choices[].delta.content // empty' |
-  awk '{ORS=""; print; fflush()} END {print "\n"}'
\ No newline at end of file
+  awk '{ORS=""; print; fflush()} END {print "\n"}'
diff --git a/uv.lock b/uv.lock
index deabdc7b..6960924b 100644
--- a/uv.lock
+++ b/uv.lock
@@ -379,7 +379,7 @@ requires-dist = [
     { name = "aiofiles", specifier = ">=24.1.0" },
     { name = "aiohttp", specifier = ">=3.12.14" },
     { name = "aiosqlite", specifier = ">=0.21.0" },
-    { name = "anyio", specifier = ">=4.10.0" },
+    { name = "anyio", specifier = ">=4.11.0" },
     { name = "base58", specifier = ">=2.1.1" },
     { name = "bidict", specifier = ">=0.23.1" },
     { name = "cobs", specifier = ">=1.2.2" },

← 9058b117 pipeline parallel fix  ·  back to Exo  ·  five billion percent better shutdown handling 364087b9 →