[object Object]

← back to Exo

memory tidy (#1558)

c45ff9ad4307341d0162d2ddbcdad50baefe1799 · 2026-02-19 21:15:33 +0000 · Evan Quiney

add some pythonic extensions to memory, did a bunch of cleanup.

Files touched

Diff

commit c45ff9ad4307341d0162d2ddbcdad50baefe1799
Author: Evan Quiney <evanev7@gmail.com>
Date:   Thu Feb 19 21:15:33 2026 +0000

    memory tidy (#1558)
    
    add some pythonic extensions to memory, did a bunch of cleanup.
---
 dashboard/src/lib/stores/app.svelte.ts             |   5 +
 dashboard/src/routes/+page.svelte                  |  11 +-
 dashboard/src/routes/downloads/+page.svelte        |  16 +--
 src/exo/download/coordinator.py                    |   8 +-
 src/exo/download/download_utils.py                 |  35 +++---
 src/exo/download/shard_downloader.py               |   6 +-
 src/exo/master/api.py                              |   2 +-
 src/exo/master/placement_utils.py                  |  13 ++-
 src/exo/master/tests/test_placement.py             |   6 +-
 .../tests/test_apply/test_apply_node_download.py   |   6 +-
 src/exo/shared/types/memory.py                     | 120 +++++++++++++++++----
 src/exo/shared/types/worker/downloads.py           |  14 +--
 src/exo/worker/engines/image/generate.py           |   4 +-
 src/exo/worker/engines/mlx/cache.py                |   2 +-
 src/exo/worker/engines/mlx/utils_mlx.py            |  21 ++--
 .../test_plan/test_download_and_loading.py         |  24 ++---
 16 files changed, 174 insertions(+), 119 deletions(-)

diff --git a/dashboard/src/lib/stores/app.svelte.ts b/dashboard/src/lib/stores/app.svelte.ts
index 3e0074e3..03379cee 100644
--- a/dashboard/src/lib/stores/app.svelte.ts
+++ b/dashboard/src/lib/stores/app.svelte.ts
@@ -250,6 +250,11 @@ interface RawStateResponse {
   >;
   // Thunderbolt bridge cycles (nodes with bridge enabled forming loops)
   thunderboltBridgeCycles?: string[][];
+  // Disk usage per node
+  nodeDisk?: Record<
+    string,
+    { total: { inBytes: number }; available: { inBytes: number } }
+  >;
 }
 
 export interface MessageAttachment {
diff --git a/dashboard/src/routes/+page.svelte b/dashboard/src/routes/+page.svelte
index 21e71774..76a3dbfd 100644
--- a/dashboard/src/routes/+page.svelte
+++ b/dashboard/src/routes/+page.svelte
@@ -858,10 +858,8 @@
     if (!progress || typeof progress !== "object") return null;
 
     const prog = progress as Record<string, unknown>;
-    const totalBytes = getBytes(prog.total_bytes ?? prog.totalBytes);
-    const downloadedBytes = getBytes(
-      prog.downloaded_bytes ?? prog.downloadedBytes,
-    );
+    const totalBytes = getBytes(prog.total);
+    const downloadedBytes = getBytes(prog.downloaded);
     const speed = (prog.speed as number) ?? 0;
     const completedFiles =
       (prog.completed_files as number) ?? (prog.completedFiles as number) ?? 0;
@@ -874,8 +872,8 @@
     for (const [fileName, fileData] of Object.entries(filesObj)) {
       if (!fileData || typeof fileData !== "object") continue;
       const fd = fileData as Record<string, unknown>;
-      const fTotal = getBytes(fd.total_bytes ?? fd.totalBytes);
-      const fDownloaded = getBytes(fd.downloaded_bytes ?? fd.downloadedBytes);
+      const fTotal = getBytes(fd.total);
+      const fDownloaded = getBytes(fd.downloaded);
       files.push({
         name: fileName,
         totalBytes: fTotal,
@@ -1264,7 +1262,6 @@
     if (typeof value === "number") return value;
     if (value && typeof value === "object") {
       const v = value as Record<string, unknown>;
-      if (typeof v.in_bytes === "number") return v.in_bytes;
       if (typeof v.inBytes === "number") return v.inBytes;
     }
     return 0;
diff --git a/dashboard/src/routes/downloads/+page.svelte b/dashboard/src/routes/downloads/+page.svelte
index 91719dc1..a57a18d3 100644
--- a/dashboard/src/routes/downloads/+page.svelte
+++ b/dashboard/src/routes/downloads/+page.svelte
@@ -74,7 +74,6 @@
     if (typeof value === "number") return value;
     if (value && typeof value === "object") {
       const v = value as Record<string, unknown>;
-      if (typeof v.in_bytes === "number") return v.in_bytes;
       if (typeof v.inBytes === "number") return v.inBytes;
     }
     return 0;
@@ -231,23 +230,14 @@
             undefined;
           let cell: CellStatus;
           if (tag === "DownloadCompleted") {
-            const totalBytes = getBytes(
-              payload.total_bytes ?? payload.totalBytes,
-            );
+            const totalBytes = getBytes(payload.total);
             cell = { kind: "completed", totalBytes, modelDirectory };
           } else if (tag === "DownloadOngoing") {
             const rawProgress =
               payload.download_progress ?? payload.downloadProgress ?? {};
             const prog = rawProgress as Record<string, unknown>;
-            const totalBytes = getBytes(
-              prog.total_bytes ??
-                prog.totalBytes ??
-                payload.total_bytes ??
-                payload.totalBytes,
-            );
-            const downloadedBytes = getBytes(
-              prog.downloaded_bytes ?? prog.downloadedBytes,
-            );
+            const totalBytes = getBytes(prog.total ?? payload.total);
+            const downloadedBytes = getBytes(prog.downloaded);
             const speed = (prog.speed as number) ?? 0;
             const etaMs =
               (prog.eta_ms as number) ?? (prog.etaMs as number) ?? 0;
diff --git a/src/exo/download/coordinator.py b/src/exo/download/coordinator.py
index 30e45a08..f2b44495 100644
--- a/src/exo/download/coordinator.py
+++ b/src/exo/download/coordinator.py
@@ -80,7 +80,7 @@ class DownloadCoordinator:
             completed = DownloadCompleted(
                 shard_metadata=callback_shard,
                 node_id=self.node_id,
-                total_bytes=progress.total_bytes,
+                total=progress.total,
                 model_directory=self._model_dir(model_id),
             )
             self.download_status[model_id] = completed
@@ -203,7 +203,7 @@ class DownloadCoordinator:
             completed = DownloadCompleted(
                 shard_metadata=shard,
                 node_id=self.node_id,
-                total_bytes=initial_progress.total_bytes,
+                total=initial_progress.total,
                 model_directory=self._model_dir(model_id),
             )
             self.download_status[model_id] = completed
@@ -332,13 +332,13 @@ class DownloadCoordinator:
                         status: DownloadProgress = DownloadCompleted(
                             node_id=self.node_id,
                             shard_metadata=progress.shard,
-                            total_bytes=progress.total_bytes,
+                            total=progress.total,
                             model_directory=self._model_dir(
                                 progress.shard.model_card.model_id
                             ),
                         )
                     elif progress.status in ["in_progress", "not_started"]:
-                        if progress.downloaded_bytes_this_session.in_bytes == 0:
+                        if progress.downloaded_this_session.in_bytes == 0:
                             status = DownloadPending(
                                 node_id=self.node_id,
                                 shard_metadata=progress.shard,
diff --git a/src/exo/download/download_utils.py b/src/exo/download/download_utils.py
index 5691e5dd..868bbbe9 100644
--- a/src/exo/download/download_utils.py
+++ b/src/exo/download/download_utils.py
@@ -80,9 +80,9 @@ def map_repo_file_download_progress_to_download_progress_data(
     repo_file_download_progress: RepoFileDownloadProgress,
 ) -> DownloadProgressData:
     return DownloadProgressData(
-        downloaded_bytes=repo_file_download_progress.downloaded,
-        downloaded_bytes_this_session=repo_file_download_progress.downloaded_this_session,
-        total_bytes=repo_file_download_progress.total,
+        downloaded=repo_file_download_progress.downloaded,
+        downloaded_this_session=repo_file_download_progress.downloaded_this_session,
+        total=repo_file_download_progress.total,
         completed_files=1 if repo_file_download_progress.status == "complete" else 0,
         total_files=1,
         speed=repo_file_download_progress.speed,
@@ -95,9 +95,9 @@ def map_repo_download_progress_to_download_progress_data(
     repo_download_progress: RepoDownloadProgress,
 ) -> DownloadProgressData:
     return DownloadProgressData(
-        total_bytes=repo_download_progress.total_bytes,
-        downloaded_bytes=repo_download_progress.downloaded_bytes,
-        downloaded_bytes_this_session=repo_download_progress.downloaded_bytes_this_session,
+        total=repo_download_progress.total,
+        downloaded=repo_download_progress.downloaded,
+        downloaded_this_session=repo_download_progress.downloaded_this_session,
         completed_files=repo_download_progress.completed_files,
         total_files=repo_download_progress.total_files,
         speed=repo_download_progress.overall_speed,
@@ -578,19 +578,20 @@ def calculate_repo_progress(
     file_progress: dict[str, RepoFileDownloadProgress],
     all_start_time: float,
 ) -> RepoDownloadProgress:
-    all_total_bytes = sum((p.total.in_bytes for p in file_progress.values()), 0)
-    all_downloaded_bytes = sum(
-        (p.downloaded.in_bytes for p in file_progress.values()), 0
+    all_total = sum((p.total for p in file_progress.values()), Memory.from_bytes(0))
+    all_downloaded = sum(
+        (p.downloaded for p in file_progress.values()), Memory.from_bytes(0)
     )
-    all_downloaded_bytes_this_session = sum(
-        (p.downloaded_this_session.in_bytes for p in file_progress.values()), 0
+    all_downloaded_this_session = sum(
+        (p.downloaded_this_session for p in file_progress.values()),
+        Memory.from_bytes(0),
     )
     elapsed_time = time.time() - all_start_time
     all_speed = (
-        all_downloaded_bytes_this_session / elapsed_time if elapsed_time > 0 else 0
+        all_downloaded_this_session.in_bytes / elapsed_time if elapsed_time > 0 else 0
     )
     all_eta = (
-        timedelta(seconds=(all_total_bytes - all_downloaded_bytes) / all_speed)
+        timedelta(seconds=(all_total - all_downloaded).in_bytes / all_speed)
         if all_speed > 0
         else timedelta(seconds=0)
     )
@@ -609,11 +610,9 @@ def calculate_repo_progress(
             [p for p in file_progress.values() if p.downloaded == p.total]
         ),
         total_files=len(file_progress),
-        downloaded_bytes=Memory.from_bytes(all_downloaded_bytes),
-        downloaded_bytes_this_session=Memory.from_bytes(
-            all_downloaded_bytes_this_session
-        ),
-        total_bytes=Memory.from_bytes(all_total_bytes),
+        downloaded=all_downloaded,
+        downloaded_this_session=all_downloaded_this_session,
+        total=all_total,
         overall_speed=all_speed,
         overall_eta=all_eta,
         status=status,
diff --git a/src/exo/download/shard_downloader.py b/src/exo/download/shard_downloader.py
index 9dd8c324..22e85643 100644
--- a/src/exo/download/shard_downloader.py
+++ b/src/exo/download/shard_downloader.py
@@ -107,9 +107,9 @@ NOOP_DOWNLOAD_PROGRESS = RepoDownloadProgress(
     ),
     completed_files=0,
     total_files=0,
-    downloaded_bytes=Memory.from_bytes(0),
-    downloaded_bytes_this_session=Memory.from_bytes(0),
-    total_bytes=Memory.from_bytes(0),
+    downloaded=Memory.from_bytes(0),
+    downloaded_this_session=Memory.from_bytes(0),
+    total=Memory.from_bytes(0),
     overall_speed=0,
     overall_eta=timedelta(seconds=0),
     status="complete",
diff --git a/src/exo/master/api.py b/src/exo/master/api.py
index 319a8164..0f3d8711 100644
--- a/src/exo/master/api.py
+++ b/src/exo/master/api.py
@@ -1322,7 +1322,7 @@ class API:
                     name=card.model_id.short(),
                     description="",
                     tags=[],
-                    storage_size_megabytes=int(card.storage_size.in_mb),
+                    storage_size_megabytes=card.storage_size.in_mb,
                     supports_tensor=card.supports_tensor,
                     tasks=[task.value for task in card.tasks],
                     is_custom=is_custom_card(card.model_id),
diff --git a/src/exo/master/placement_utils.py b/src/exo/master/placement_utils.py
index b20a39cc..b80d70c0 100644
--- a/src/exo/master/placement_utils.py
+++ b/src/exo/master/placement_utils.py
@@ -102,22 +102,21 @@ def _allocate_and_validate_layers(
     layer_allocations = allocate_layers_proportionally(
         total_layers=model_card.n_layers,
         memory_fractions=[
-            node_memory[node_id].ram_available.in_bytes / total_memory.in_bytes
-            for node_id in node_ids
+            node_memory[node_id].ram_available / total_memory for node_id in node_ids
         ],
     )
 
-    total_storage_bytes = model_card.storage_size.in_bytes
+    total_storage = model_card.storage_size
     total_layers = model_card.n_layers
     for i, node_id in enumerate(node_ids):
         node_layers = layer_allocations[i]
-        required_memory = (total_storage_bytes * node_layers) // total_layers
-        available_memory = node_memory[node_id].ram_available.in_bytes
+        required_memory = (total_storage * node_layers) // total_layers
+        available_memory = node_memory[node_id].ram_available
         if required_memory > available_memory:
             raise ValueError(
                 f"Node {i} ({node_id}) has insufficient memory: "
-                f"requires {required_memory / (1024**3):.2f} GB for {node_layers} layers, "
-                f"but only has {available_memory / (1024**3):.2f} GB available"
+                f"requires {required_memory.in_gb:.2f} GB for {node_layers} layers, "
+                f"but only has {available_memory.in_gb:.2f} GB available"
             )
 
     return layer_allocations
diff --git a/src/exo/master/tests/test_placement.py b/src/exo/master/tests/test_placement.py
index ad5638e7..cad495ea 100644
--- a/src/exo/master/tests/test_placement.py
+++ b/src/exo/master/tests/test_placement.py
@@ -80,8 +80,8 @@ def test_get_instance_placements_create_instance(
 ):
     # arrange
     model_card.n_layers = total_layers
-    model_card.storage_size.in_bytes = sum(
-        available_memory
+    model_card.storage_size = Memory.from_bytes(
+        sum(available_memory)
     )  # make it exactly fit across all nodes
     topology = Topology()
 
@@ -349,7 +349,7 @@ def test_tensor_rdma_backend_connectivity_matrix(
     # arrange
     topology = Topology()
     model_card.n_layers = 12
-    model_card.storage_size.in_bytes = 1500
+    model_card.storage_size = Memory.from_bytes(1500)
 
     node_a = NodeId()
     node_b = NodeId()
diff --git a/src/exo/shared/tests/test_apply/test_apply_node_download.py b/src/exo/shared/tests/test_apply/test_apply_node_download.py
index 6b1cb8cc..f9df6e07 100644
--- a/src/exo/shared/tests/test_apply/test_apply_node_download.py
+++ b/src/exo/shared/tests/test_apply/test_apply_node_download.py
@@ -14,7 +14,7 @@ def test_apply_node_download_progress():
     event = DownloadCompleted(
         node_id=NodeId("node-1"),
         shard_metadata=shard1,
-        total_bytes=Memory(),
+        total=Memory(),
     )
 
     new_state = apply_node_download_progress(
@@ -30,12 +30,12 @@ def test_apply_two_node_download_progress():
     event1 = DownloadCompleted(
         node_id=NodeId("node-1"),
         shard_metadata=shard1,
-        total_bytes=Memory(),
+        total=Memory(),
     )
     event2 = DownloadCompleted(
         node_id=NodeId("node-1"),
         shard_metadata=shard2,
-        total_bytes=Memory(),
+        total=Memory(),
     )
     state = State(downloads={NodeId("node-1"): [event1]})
 
diff --git a/src/exo/shared/types/memory.py b/src/exo/shared/types/memory.py
index b97fb345..2684d9b9 100644
--- a/src/exo/shared/types/memory.py
+++ b/src/exo/shared/types/memory.py
@@ -1,10 +1,10 @@
 from math import ceil
-from typing import Self
+from typing import Self, overload
 
-from exo.utils.pydantic_ext import CamelCaseModel
+from exo.utils.pydantic_ext import FrozenModel
 
 
-class Memory(CamelCaseModel):
+class Memory(FrozenModel):
     in_bytes: int = 0
 
     @classmethod
@@ -33,12 +33,22 @@ class Memory(CamelCaseModel):
         return cls(in_bytes=round(val * 1024))
 
     @property
-    def in_mb(self) -> float:
-        """The approximate megabytes this memory represents. Setting this property rounds to the nearest byte."""
-        return self.in_bytes / (1024**2)
+    def in_mb(self) -> int:
+        """The approximate megabytes this memory represents, rounded to nearest MB. Setting this property rounds to the nearest byte."""
+        return round(self.in_bytes / (1024**2))
 
     @in_mb.setter
-    def in_mb(self, val: float):
+    def in_mb(self, val: int):
+        """Set the megabytes for this memory."""
+        self.in_bytes = val * (1024**2)
+
+    @property
+    def in_float_mb(self) -> float:
+        """The megabytes this memory represents as a float. Setting this property rounds to the nearest byte."""
+        return self.in_bytes / (1024**2)
+
+    @in_float_mb.setter
+    def in_float_mb(self, val: float):
         """Set the megabytes for this memory, rounded to the nearest byte."""
         self.in_bytes = round(val * (1024**2))
 
@@ -57,17 +67,85 @@ class Memory(CamelCaseModel):
         """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)
-
-    def __lt__(self, other: Self) -> bool:
-        return self.in_bytes < other.in_bytes
-
-    def __le__(self, other: Self) -> bool:
-        return self.in_bytes <= other.in_bytes
-
-    def __gt__(self, other: Self) -> bool:
-        return self.in_bytes > other.in_bytes
-
-    def __ge__(self, other: Self) -> bool:
-        return self.in_bytes >= other.in_bytes
+    def __add__(self, other: object) -> "Memory":
+        if isinstance(other, Memory):
+            return Memory.from_bytes(self.in_bytes + other.in_bytes)
+        return NotImplemented
+
+    def __radd__(self, other: object) -> "Memory":
+        if other == 0:
+            return self
+        return NotImplemented
+
+    def __sub__(self, other: object) -> "Memory":
+        if isinstance(other, Memory):
+            return Memory.from_bytes(self.in_bytes - other.in_bytes)
+        return NotImplemented
+
+    def __mul__(self, other: int | float):
+        return Memory.from_bytes(round(self.in_bytes * other))
+
+    def __rmul__(self, other: int | float):
+        return self * other
+
+    @overload
+    def __truediv__(self, other: "Memory") -> float: ...
+    @overload
+    def __truediv__(self, other: int) -> "Memory": ...
+    @overload
+    def __truediv__(self, other: float) -> "Memory": ...
+    def __truediv__(self, other: object) -> "Memory | float":
+        if isinstance(other, Memory):
+            return self.in_bytes / other.in_bytes
+        if isinstance(other, (int, float)):
+            return Memory.from_bytes(round(self.in_bytes / other))
+        return NotImplemented
+
+    def __floordiv__(self, other: object) -> "Memory":
+        if isinstance(other, (int, float)):
+            return Memory.from_bytes(int(self.in_bytes // other))
+        return NotImplemented
+
+    def __lt__(self, other: object) -> bool:
+        if isinstance(other, Memory):
+            return self.in_bytes < other.in_bytes
+        return NotImplemented
+
+    def __le__(self, other: object) -> bool:
+        if isinstance(other, Memory):
+            return self.in_bytes <= other.in_bytes
+        return NotImplemented
+
+    def __gt__(self, other: object) -> bool:
+        if isinstance(other, Memory):
+            return self.in_bytes > other.in_bytes
+        return NotImplemented
+
+    def __ge__(self, other: object) -> bool:
+        if isinstance(other, Memory):
+            return self.in_bytes >= other.in_bytes
+        return NotImplemented
+
+    def __eq__(self, other: object) -> bool:
+        if isinstance(other, Memory):
+            return self.in_bytes == other.in_bytes
+        return NotImplemented
+
+    def __repr__(self) -> str:
+        return f"Memory.from_bytes({self.in_bytes})"
+
+    def __str__(self) -> str:
+        if self.in_gb > 2:
+            val = self.in_gb
+            unit = "GiB"
+        elif self.in_mb > 2:
+            val = self.in_mb
+            unit = "MiB"
+        elif self.in_kb > 3:
+            val = self.in_kb
+            unit = "KiB"
+        else:
+            val = self.in_bytes
+            unit = "B"
+
+        return f"{val:.2f} {unit}".rstrip("0").rstrip(".") + f" {unit}"
diff --git a/src/exo/shared/types/worker/downloads.py b/src/exo/shared/types/worker/downloads.py
index a29edcbf..45762846 100644
--- a/src/exo/shared/types/worker/downloads.py
+++ b/src/exo/shared/types/worker/downloads.py
@@ -10,9 +10,9 @@ from exo.utils.pydantic_ext import CamelCaseModel, TaggedModel
 
 
 class DownloadProgressData(CamelCaseModel):
-    total_bytes: Memory
-    downloaded_bytes: Memory
-    downloaded_bytes_this_session: Memory
+    total: Memory
+    downloaded: Memory
+    downloaded_this_session: Memory
 
     completed_files: int
     total_files: int
@@ -34,7 +34,7 @@ class DownloadPending(BaseDownloadProgress):
 
 
 class DownloadCompleted(BaseDownloadProgress):
-    total_bytes: Memory
+    total: Memory
 
 
 class DownloadFailed(BaseDownloadProgress):
@@ -86,9 +86,9 @@ class RepoDownloadProgress(BaseModel):
     shard: ShardMetadata
     completed_files: int
     total_files: int
-    downloaded_bytes: Memory
-    downloaded_bytes_this_session: Memory
-    total_bytes: Memory
+    downloaded: Memory
+    downloaded_this_session: Memory
+    total: Memory
     overall_speed: float
     overall_eta: timedelta
     status: Literal["not_started", "in_progress", "complete"]
diff --git a/src/exo/worker/engines/image/generate.py b/src/exo/worker/engines/image/generate.py
index f5526c8f..2f4c5a4e 100644
--- a/src/exo/worker/engines/image/generate.py
+++ b/src/exo/worker/engines/image/generate.py
@@ -166,7 +166,7 @@ def generate_image(
                             else 0.0
                         )
 
-                        peak_memory_gb = mx.get_peak_memory() / (1024**3)
+                        peak_memory = Memory.from_bytes(mx.get_peak_memory())
 
                         stats = ImageGenerationStats(
                             seconds_per_step=seconds_per_step,
@@ -175,7 +175,7 @@ def generate_image(
                             num_images=num_images,
                             image_width=width,
                             image_height=height,
-                            peak_memory_usage=Memory.from_gb(peak_memory_gb),
+                            peak_memory_usage=peak_memory,
                         )
 
                     buffer = io.BytesIO()
diff --git a/src/exo/worker/engines/mlx/cache.py b/src/exo/worker/engines/mlx/cache.py
index 7669f1c1..ae6f76fa 100644
--- a/src/exo/worker/engines/mlx/cache.py
+++ b/src/exo/worker/engines/mlx/cache.py
@@ -22,7 +22,7 @@ from exo.worker.runner.bootstrap import logger
 # Fraction of device memory above which LRU eviction kicks in.
 # Smaller machines need more aggressive eviction.
 def _default_memory_threshold() -> float:
-    total_gb = psutil.virtual_memory().total / (1024**3)
+    total_gb = Memory.from_bytes(psutil.virtual_memory().total).in_gb
     if total_gb >= 128:
         return 0.85
     if total_gb >= 64:
diff --git a/src/exo/worker/engines/mlx/utils_mlx.py b/src/exo/worker/engines/mlx/utils_mlx.py
index 01bfbe50..360a018b 100644
--- a/src/exo/worker/engines/mlx/utils_mlx.py
+++ b/src/exo/worker/engines/mlx/utils_mlx.py
@@ -232,11 +232,11 @@ def shard_and_load(
 
     # Estimate timeout based on model size (5x default for large queued workloads)
     base_timeout = float(os.environ.get("EXO_MODEL_LOAD_TIMEOUT", "300"))
-    model_size_gb = get_weights_size(shard_metadata).in_bytes / (1024**3)
-    timeout_seconds = base_timeout + model_size_gb
+    model_size = get_weights_size(shard_metadata)
+    timeout_seconds = base_timeout + model_size.in_gb
     logger.info(
         f"Evaluating model parameters with timeout of {timeout_seconds:.0f}s "
-        f"(model size: {model_size_gb:.1f}GB)"
+        f"(model size: {model_size.in_gb:.1f}GB)"
     )
 
     match shard_metadata:
@@ -642,18 +642,17 @@ def set_wired_limit_for_model(model_size: Memory):
     if not mx.metal.is_available():
         return
 
-    model_bytes = model_size.in_bytes
-    max_rec_size = int(mx.metal.device_info()["max_recommended_working_set_size"])
-    if model_bytes > 0.9 * max_rec_size:
-        model_mb = model_bytes // 2**20
-        max_rec_mb = max_rec_size // 2**20
+    max_rec_size = Memory.from_bytes(
+        int(mx.metal.device_info()["max_recommended_working_set_size"])
+    )
+    if model_size > 0.9 * max_rec_size:
         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} "
+            f"Generating with a model that requires {model_size.in_float_mb:.1f} MB "
+            f"which is close to the maximum recommended size of {max_rec_size.in_float_mb:.1f} "
             "MB. This can be slow. See the documentation for possible work-arounds: "
             "https://github.com/ml-explore/mlx-lm/tree/main#large-models"
         )
-    mx.set_wired_limit(max_rec_size)
+    mx.set_wired_limit(max_rec_size.in_bytes)
     logger.info(f"Wired limit set to {max_rec_size}.")
 
 
diff --git a/src/exo/worker/tests/unittests/test_plan/test_download_and_loading.py b/src/exo/worker/tests/unittests/test_plan/test_download_and_loading.py
index 9c318517..abcb4939 100644
--- a/src/exo/worker/tests/unittests/test_plan/test_download_and_loading.py
+++ b/src/exo/worker/tests/unittests/test_plan/test_download_and_loading.py
@@ -90,14 +90,10 @@ def test_plan_loads_model_when_all_shards_downloaded_and_waiting():
 
     global_download_status = {
         NODE_A: [
-            DownloadCompleted(
-                shard_metadata=shard1, node_id=NODE_A, total_bytes=Memory()
-            )
+            DownloadCompleted(shard_metadata=shard1, node_id=NODE_A, total=Memory())
         ],
         NODE_B: [
-            DownloadCompleted(
-                shard_metadata=shard2, node_id=NODE_B, total_bytes=Memory()
-            )
+            DownloadCompleted(shard_metadata=shard2, node_id=NODE_B, total=Memory())
         ],
     }
 
@@ -138,9 +134,7 @@ def test_plan_does_not_request_download_when_shard_already_downloaded():
     # Global state shows shard is downloaded for NODE_A
     global_download_status: dict[NodeId, list[DownloadProgress]] = {
         NODE_A: [
-            DownloadCompleted(
-                shard_metadata=shard, node_id=NODE_A, total_bytes=Memory()
-            )
+            DownloadCompleted(shard_metadata=shard, node_id=NODE_A, total=Memory())
         ],
         NODE_B: [],
     }
@@ -187,9 +181,7 @@ def test_plan_does_not_load_model_until_all_shards_downloaded_globally():
 
     global_download_status = {
         NODE_A: [
-            DownloadCompleted(
-                shard_metadata=shard1, node_id=NODE_A, total_bytes=Memory()
-            )
+            DownloadCompleted(shard_metadata=shard1, node_id=NODE_A, total=Memory())
         ],
         NODE_B: [],  # NODE_B has no downloads completed yet
     }
@@ -207,14 +199,10 @@ def test_plan_does_not_load_model_until_all_shards_downloaded_globally():
 
     global_download_status = {
         NODE_A: [
-            DownloadCompleted(
-                shard_metadata=shard1, node_id=NODE_A, total_bytes=Memory()
-            )
+            DownloadCompleted(shard_metadata=shard1, node_id=NODE_A, total=Memory())
         ],
         NODE_B: [
-            DownloadCompleted(
-                shard_metadata=shard2, node_id=NODE_B, total_bytes=Memory()
-            )
+            DownloadCompleted(shard_metadata=shard2, node_id=NODE_B, total=Memory())
         ],  # NODE_B has no downloads completed yet
     }
 

← 7031901a Prevent common fatal crashes (#1555)  ·  back to Exo  ·  Prioritise tb for ring instances (#1556) f662c129 →