← back to Exo
Reduce reliance on internet (#1363)
a0f4f363555744f2a9660679437be064bb2bb712 · 2026-02-03 20:03:29 +0000 · rltakashige
## Motivation
Offline users currently have to wait for every retry to fail before
being able to launch a model.
For users that restart clusters often or share API keys between devices,
we also spam HuggingFace with downloads every 5 minutes.
These issues are caused by _emit_existing_download_progress being
inefficient.
## Changes
- Only query HuggingFace once while EXO is running (assumption being
that a change should only be reflected on a new EXO session)
- Only query HuggingFace when there is an internet connection (polling
connectivity every 10 seconds)
- Request download progress if we switch from no connectivity ->
connected to reduce the wait.
- Reduce download progress sleep as it's no longer expensive (queries
cache most of the time).
- Reduce retries as 30 is way too many.
## Test Plan
### Manual Testing
Manually tested the behaviour.
### Automated Testing
None, should I add any? We do have some tests for this folder, but they
are probably not too helpful.
Files touched
M src/exo/download/coordinator.pyM src/exo/download/download_utils.pyM src/exo/download/impl_shard_downloader.pyM src/exo/download/shard_downloader.pyM src/exo/shared/models/model_cards.py
Diff
commit a0f4f363555744f2a9660679437be064bb2bb712
Author: rltakashige <rl.takashige@gmail.com>
Date: Tue Feb 3 20:03:29 2026 +0000
Reduce reliance on internet (#1363)
## Motivation
Offline users currently have to wait for every retry to fail before
being able to launch a model.
For users that restart clusters often or share API keys between devices,
we also spam HuggingFace with downloads every 5 minutes.
These issues are caused by _emit_existing_download_progress being
inefficient.
## Changes
- Only query HuggingFace once while EXO is running (assumption being
that a change should only be reflected on a new EXO session)
- Only query HuggingFace when there is an internet connection (polling
connectivity every 10 seconds)
- Request download progress if we switch from no connectivity ->
connected to reduce the wait.
- Reduce download progress sleep as it's no longer expensive (queries
cache most of the time).
- Reduce retries as 30 is way too many.
## Test Plan
### Manual Testing
Manually tested the behaviour.
### Automated Testing
None, should I add any? We do have some tests for this folder, but they
are probably not too helpful.
---
src/exo/download/coordinator.py | 34 +++++++++++--
src/exo/download/download_utils.py | 82 ++++++++++++++++++++++++-------
src/exo/download/impl_shard_downloader.py | 28 +++++++++--
src/exo/download/shard_downloader.py | 5 ++
src/exo/shared/models/model_cards.py | 8 +--
5 files changed, 130 insertions(+), 27 deletions(-)
diff --git a/src/exo/download/coordinator.py b/src/exo/download/coordinator.py
index c2f7b9e9..f5798ad3 100644
--- a/src/exo/download/coordinator.py
+++ b/src/exo/download/coordinator.py
@@ -1,4 +1,5 @@
import asyncio
+import socket
from dataclasses import dataclass, field
from typing import Iterator
@@ -60,10 +61,37 @@ class DownloadCoordinator:
async def run(self) -> None:
logger.info("Starting DownloadCoordinator")
+ self._test_internet_connection()
async with self._tg as tg:
tg.start_soon(self._command_processor)
tg.start_soon(self._forward_events)
tg.start_soon(self._emit_existing_download_progress)
+ tg.start_soon(self._check_internet_connection)
+
+ def _test_internet_connection(self) -> None:
+ try:
+ socket.create_connection(("1.1.1.1", 443), timeout=3).close()
+ self.shard_downloader.set_internet_connection(True)
+ except OSError:
+ self.shard_downloader.set_internet_connection(False)
+ logger.debug(
+ f"Internet connectivity: {self.shard_downloader.internet_connection}"
+ )
+
+ async def _check_internet_connection(self) -> None:
+ first_connection = True
+ while True:
+ await asyncio.sleep(10)
+
+ # Assume that internet connection is set to False on 443 errors.
+ if self.shard_downloader.internet_connection:
+ continue
+
+ self._test_internet_connection()
+
+ if first_connection and self.shard_downloader.internet_connection:
+ first_connection = False
+ self._tg.start_soon(self._emit_existing_download_progress)
def shutdown(self) -> None:
self._tg.cancel_scope.cancel()
@@ -241,7 +269,7 @@ class DownloadCoordinator:
async def _emit_existing_download_progress(self) -> None:
try:
while True:
- logger.info(
+ logger.debug(
"DownloadCoordinator: Fetching and emitting existing download progress..."
)
async for (
@@ -274,10 +302,10 @@ class DownloadCoordinator:
await self.event_sender.send(
NodeDownloadProgress(download_progress=status)
)
- logger.info(
+ logger.debug(
"DownloadCoordinator: Done emitting existing download progress."
)
- await anyio.sleep(5 * 60) # 5 minutes
+ await anyio.sleep(60)
except Exception as e:
logger.error(
f"DownloadCoordinator: Error emitting existing download progress: {e}"
diff --git a/src/exo/download/download_utils.py b/src/exo/download/download_utils.py
index 6dec4718..618e4f38 100644
--- a/src/exo/download/download_utils.py
+++ b/src/exo/download/download_utils.py
@@ -49,6 +49,10 @@ class HuggingFaceAuthenticationError(Exception):
"""Raised when HuggingFace returns 401/403 for a model download."""
+class HuggingFaceRateLimitError(Exception):
+ """429 Huggingface code"""
+
+
async def _build_auth_error_message(status_code: int, model_id: ModelId) -> str:
token = await get_hf_token()
if status_code == 401 and token is None:
@@ -154,49 +158,76 @@ async def seed_models(seed_dir: str | Path):
logger.error(traceback.format_exc())
+_fetched_file_lists_this_session: set[str] = set()
+
+
async def fetch_file_list_with_cache(
- model_id: ModelId, revision: str = "main", recursive: bool = False
+ model_id: ModelId,
+ revision: str = "main",
+ recursive: bool = False,
+ skip_internet: bool = False,
+ on_connection_lost: Callable[[], None] = lambda: None,
) -> list[FileListEntry]:
target_dir = (await ensure_models_dir()) / "caches" / model_id.normalize()
await aios.makedirs(target_dir, exist_ok=True)
cache_file = target_dir / f"{model_id.normalize()}--{revision}--file_list.json"
+ cache_key = f"{model_id.normalize()}--{revision}"
+
+ if cache_key in _fetched_file_lists_this_session and await aios.path.exists(
+ cache_file
+ ):
+ async with aiofiles.open(cache_file, "r") as f:
+ return TypeAdapter(list[FileListEntry]).validate_json(await f.read())
+
+ if skip_internet:
+ if await aios.path.exists(cache_file):
+ async with aiofiles.open(cache_file, "r") as f:
+ return TypeAdapter(list[FileListEntry]).validate_json(await f.read())
+ raise FileNotFoundError(
+ f"No internet connection and no cached file list for {model_id}"
+ )
- # Always try fresh first
try:
file_list = await fetch_file_list_with_retry(
- model_id, revision, recursive=recursive
+ model_id,
+ revision,
+ recursive=recursive,
+ on_connection_lost=on_connection_lost,
)
- # Update cache with fresh data
async with aiofiles.open(cache_file, "w") as f:
await f.write(
TypeAdapter(list[FileListEntry]).dump_json(file_list).decode()
)
+ _fetched_file_lists_this_session.add(cache_key)
return file_list
except Exception as e:
- # Fetch failed - try cache fallback
if await aios.path.exists(cache_file):
logger.warning(
f"Failed to fetch file list for {model_id}, using cached data: {e}"
)
async with aiofiles.open(cache_file, "r") as f:
return TypeAdapter(list[FileListEntry]).validate_json(await f.read())
- # No cache available, propagate the error
- raise
+ raise FileNotFoundError(f"Failed to fetch file list for {model_id}: {e}") from e
async def fetch_file_list_with_retry(
- model_id: ModelId, revision: str = "main", path: str = "", recursive: bool = False
+ model_id: ModelId,
+ revision: str = "main",
+ path: str = "",
+ recursive: bool = False,
+ on_connection_lost: Callable[[], None] = lambda: None,
) -> list[FileListEntry]:
- n_attempts = 30
+ n_attempts = 3
for attempt in range(n_attempts):
try:
return await _fetch_file_list(model_id, revision, path, recursive)
except HuggingFaceAuthenticationError:
raise
except Exception as e:
+ on_connection_lost()
if attempt == n_attempts - 1:
raise e
- await asyncio.sleep(min(8, 0.1 * float(2.0 ** int(attempt))))
+ await asyncio.sleep(2.0**attempt)
raise Exception(
f"Failed to fetch file list for {model_id=} {revision=} {path=} {recursive=}"
)
@@ -216,7 +247,11 @@ async def _fetch_file_list(
if response.status in [401, 403]:
msg = await _build_auth_error_message(response.status, model_id)
raise HuggingFaceAuthenticationError(msg)
- if response.status == 200:
+ elif response.status == 429:
+ raise HuggingFaceRateLimitError(
+ f"Couldn't download {model_id} because of HuggingFace rate limit."
+ )
+ elif response.status == 200:
data_json = await response.text()
data = TypeAdapter(list[FileListEntry]).validate_json(data_json)
files: list[FileListEntry] = []
@@ -249,7 +284,7 @@ def create_http_session(
else:
total_timeout = 1800
connect_timeout = 60
- sock_read_timeout = 1800
+ sock_read_timeout = 60
sock_connect_timeout = 60
ssl_context = ssl.create_default_context(
@@ -324,8 +359,9 @@ async def download_file_with_retry(
path: str,
target_dir: Path,
on_progress: Callable[[int, int, bool], None] = lambda _, __, ___: None,
+ on_connection_lost: Callable[[], None] = lambda: None,
) -> Path:
- n_attempts = 30
+ n_attempts = 3
for attempt in range(n_attempts):
try:
return await _download_file(
@@ -333,14 +369,19 @@ async def download_file_with_retry(
)
except HuggingFaceAuthenticationError:
raise
- except Exception as e:
- if isinstance(e, FileNotFoundError) or attempt == n_attempts - 1:
+ except HuggingFaceRateLimitError as e:
+ if attempt == n_attempts - 1:
raise e
logger.error(
f"Download error on attempt {attempt}/{n_attempts} for {model_id=} {revision=} {path=} {target_dir=}"
)
logger.error(traceback.format_exc())
- await asyncio.sleep(min(8, 0.1 * (2.0**attempt)))
+ await asyncio.sleep(2.0**attempt)
+ except Exception as e:
+ on_connection_lost()
+ if attempt == n_attempts - 1:
+ raise e
+ break
raise Exception(
f"Failed to download file {model_id=} {revision=} {path=} {target_dir=}"
)
@@ -542,7 +583,9 @@ async def download_shard(
on_progress: Callable[[ShardMetadata, RepoDownloadProgress], Awaitable[None]],
max_parallel_downloads: int = 8,
skip_download: bool = False,
+ skip_internet: bool = False,
allow_patterns: list[str] | None = None,
+ on_connection_lost: Callable[[], None] = lambda: None,
) -> tuple[Path, RepoDownloadProgress]:
if not skip_download:
logger.debug(f"Downloading {shard.model_card.model_id=}")
@@ -562,7 +605,11 @@ async def download_shard(
all_start_time = time.time()
file_list = await fetch_file_list_with_cache(
- shard.model_card.model_id, revision, recursive=True
+ shard.model_card.model_id,
+ revision,
+ recursive=True,
+ skip_internet=skip_internet,
+ on_connection_lost=on_connection_lost,
)
filtered_file_list = list(
filter_repo_objects(
@@ -672,6 +719,7 @@ async def download_shard(
lambda curr_bytes, total_bytes, is_renamed: schedule_progress(
file, curr_bytes, total_bytes, is_renamed
),
+ on_connection_lost=on_connection_lost,
)
if not skip_download:
diff --git a/src/exo/download/impl_shard_downloader.py b/src/exo/download/impl_shard_downloader.py
index 1b7f5eab..0e7aea1e 100644
--- a/src/exo/download/impl_shard_downloader.py
+++ b/src/exo/download/impl_shard_downloader.py
@@ -1,4 +1,5 @@
import asyncio
+from asyncio import create_task
from collections.abc import Awaitable
from pathlib import Path
from typing import AsyncIterator, Callable
@@ -49,6 +50,10 @@ class SingletonShardDownloader(ShardDownloader):
self.shard_downloader = shard_downloader
self.active_downloads: dict[ShardMetadata, asyncio.Task[Path]] = {}
+ def set_internet_connection(self, value: bool) -> None:
+ self.internet_connection = value
+ self.shard_downloader.set_internet_connection(value)
+
def on_progress(
self,
callback: Callable[[ShardMetadata, RepoDownloadProgress], Awaitable[None]],
@@ -85,6 +90,10 @@ class CachedShardDownloader(ShardDownloader):
self.shard_downloader = shard_downloader
self.cache: dict[tuple[str, ShardMetadata], Path] = {}
+ def set_internet_connection(self, value: bool) -> None:
+ self.internet_connection = value
+ self.shard_downloader.set_internet_connection(value)
+
def on_progress(
self,
callback: Callable[[ShardMetadata, RepoDownloadProgress], Awaitable[None]],
@@ -142,6 +151,8 @@ class ResumableShardDownloader(ShardDownloader):
self.on_progress_wrapper,
max_parallel_downloads=self.max_parallel_downloads,
allow_patterns=allow_patterns,
+ skip_internet=not self.internet_connection,
+ on_connection_lost=lambda: self.set_internet_connection(False),
)
return target_dir
@@ -154,12 +165,23 @@ class ResumableShardDownloader(ShardDownloader):
"""Helper coroutine that builds the shard for a model and gets its download status."""
shard = await build_full_shard(model_id)
return await download_shard(
- shard, self.on_progress_wrapper, skip_download=True
+ shard,
+ self.on_progress_wrapper,
+ skip_download=True,
+ skip_internet=not self.internet_connection,
+ on_connection_lost=lambda: self.set_internet_connection(False),
)
- # Kick off download status coroutines concurrently
+ semaphore = asyncio.Semaphore(self.max_parallel_downloads)
+
+ async def download_with_semaphore(
+ model_card: ModelCard,
+ ) -> tuple[Path, RepoDownloadProgress]:
+ async with semaphore:
+ return await _status_for_model(model_card.model_id)
+
tasks = [
- asyncio.create_task(_status_for_model(model_card.model_id))
+ create_task(download_with_semaphore(model_card))
for model_card in await get_model_cards()
]
diff --git a/src/exo/download/shard_downloader.py b/src/exo/download/shard_downloader.py
index 30c11d25..9dd8c324 100644
--- a/src/exo/download/shard_downloader.py
+++ b/src/exo/download/shard_downloader.py
@@ -16,6 +16,11 @@ from exo.shared.types.worker.shards import (
# TODO: the PipelineShardMetadata getting reinstantiated is a bit messy. Should this be a classmethod?
class ShardDownloader(ABC):
+ internet_connection: bool = False
+
+ def set_internet_connection(self, value: bool) -> None:
+ self.internet_connection = value
+
@abstractmethod
async def ensure_shard(
self, shard: ShardMetadata, config_only: bool = False
diff --git a/src/exo/shared/models/model_cards.py b/src/exo/shared/models/model_cards.py
index bf2a892f..58d84849 100644
--- a/src/exo/shared/models/model_cards.py
+++ b/src/exo/shared/models/model_cards.py
@@ -108,9 +108,9 @@ class ModelCard(CamelCaseModel):
async def fetch_from_hf(model_id: ModelId) -> "ModelCard":
"""Fetches storage size and number of layers for a Hugging Face model, returns Pydantic ModelMeta."""
# TODO: failure if files do not exist
- config_data = await get_config_data(model_id)
+ config_data = await fetch_config_data(model_id)
num_layers = config_data.layer_count
- mem_size_bytes = await get_safetensors_size(model_id)
+ mem_size_bytes = await fetch_safetensors_size(model_id)
mc = ModelCard(
model_id=ModelId(model_id),
@@ -258,7 +258,7 @@ class ConfigData(BaseModel):
return data
-async def get_config_data(model_id: ModelId) -> ConfigData:
+async def fetch_config_data(model_id: ModelId) -> ConfigData:
"""Downloads and parses config.json for a model."""
from exo.download.download_utils import (
download_file_with_retry,
@@ -280,7 +280,7 @@ async def get_config_data(model_id: ModelId) -> ConfigData:
return ConfigData.model_validate_json(await f.read())
-async def get_safetensors_size(model_id: ModelId) -> Memory:
+async def fetch_safetensors_size(model_id: ModelId) -> Memory:
"""Gets model size from safetensors index or falls back to HF API."""
from exo.download.download_utils import (
download_file_with_retry,
← acb97127 Normalize TextGenerationTaskParams.input to list[InputMessag
·
back to Exo
·
feat: add custom HuggingFace model support (#1368) 20632789 →