1515)
1616
1717
18- def exo_shard_downloader (max_parallel_downloads : int = 8 ) -> ShardDownloader :
18+ def exo_shard_downloader (
19+ max_parallel_downloads : int = 8 , offline : bool = False
20+ ) -> ShardDownloader :
1921 return SingletonShardDownloader (
20- CachedShardDownloader (ResumableShardDownloader (max_parallel_downloads ))
22+ CachedShardDownloader (
23+ ResumableShardDownloader (max_parallel_downloads , offline = offline )
24+ )
2125 )
2226
2327
@@ -50,10 +54,6 @@ def __init__(self, shard_downloader: ShardDownloader):
5054 self .shard_downloader = shard_downloader
5155 self .active_downloads : dict [ShardMetadata , asyncio .Task [Path ]] = {}
5256
53- def set_internet_connection (self , value : bool ) -> None :
54- self .internet_connection = value
55- self .shard_downloader .set_internet_connection (value )
56-
5757 def on_progress (
5858 self ,
5959 callback : Callable [[ShardMetadata , RepoDownloadProgress ], Awaitable [None ]],
@@ -90,10 +90,6 @@ def __init__(self, shard_downloader: ShardDownloader):
9090 self .shard_downloader = shard_downloader
9191 self .cache : dict [tuple [str , ShardMetadata ], Path ] = {}
9292
93- def set_internet_connection (self , value : bool ) -> None :
94- self .internet_connection = value
95- self .shard_downloader .set_internet_connection (value )
96-
9793 def on_progress (
9894 self ,
9995 callback : Callable [[ShardMetadata , RepoDownloadProgress ], Awaitable [None ]],
@@ -123,8 +119,9 @@ async def get_shard_download_status_for_shard(
123119
124120
125121class ResumableShardDownloader (ShardDownloader ):
126- def __init__ (self , max_parallel_downloads : int = 8 ):
122+ def __init__ (self , max_parallel_downloads : int = 8 , offline : bool = False ):
127123 self .max_parallel_downloads = max_parallel_downloads
124+ self .offline = offline
128125 self .on_progress_callbacks : list [
129126 Callable [[ShardMetadata , RepoDownloadProgress ], Awaitable [None ]]
130127 ] = []
@@ -151,8 +148,7 @@ async def ensure_shard(
151148 self .on_progress_wrapper ,
152149 max_parallel_downloads = self .max_parallel_downloads ,
153150 allow_patterns = allow_patterns ,
154- skip_internet = not self .internet_connection ,
155- on_connection_lost = lambda : self .set_internet_connection (False ),
151+ skip_internet = self .offline ,
156152 )
157153 return target_dir
158154
@@ -168,8 +164,7 @@ async def _status_for_model(
168164 shard ,
169165 self .on_progress_wrapper ,
170166 skip_download = True ,
171- skip_internet = not self .internet_connection ,
172- on_connection_lost = lambda : self .set_internet_connection (False ),
167+ skip_internet = self .offline ,
173168 )
174169
175170 semaphore = asyncio .Semaphore (self .max_parallel_downloads )
@@ -198,7 +193,6 @@ async def get_shard_download_status_for_shard(
198193 shard ,
199194 self .on_progress_wrapper ,
200195 skip_download = True ,
201- skip_internet = not self .internet_connection ,
202- on_connection_lost = lambda : self .set_internet_connection (False ),
196+ skip_internet = self .offline ,
203197 )
204198 return progress
0 commit comments