From 55c4385db593f1bb57a53918e0bb557a5bd9e26a Mon Sep 17 00:00:00 2001 From: Alex Cheema Date: Thu, 30 Jan 2025 20:11:06 +0000 Subject: [PATCH 1/2] cleanup tmp files on failed download --- exo/download/new_shard_download.py | 37 ++++++++++++++++++------------ 1 file changed, 22 insertions(+), 15 deletions(-) diff --git a/exo/download/new_shard_download.py b/exo/download/new_shard_download.py index 35ae31dc..e4efb3df 100644 --- a/exo/download/new_shard_download.py +++ b/exo/download/new_shard_download.py @@ -92,20 +92,28 @@ async def fetch_file_list(repo_id, revision, path=""): @retry(stop=stop_after_attempt(5), wait=wait_exponential(multiplier=0.5)) async def download_file(repo_id: str, revision: str, path: str, target_dir: Path, on_progress: Callable[[int, int], None] = lambda _, __: None) -> Path: - if (target_dir/path).exists(): return target_dir/path - await aios.makedirs((target_dir/path).parent, exist_ok=True) - base_url = f"{get_hf_endpoint()}/{repo_id}/resolve/{revision}/" - url = urljoin(base_url, path) - headers = await get_auth_headers() - async with aiohttp.ClientSession(timeout=aiohttp.ClientTimeout(total=1800, connect=60, sock_read=1800, sock_connect=60)) as session: - async with session.get(url, headers=headers, timeout=aiohttp.ClientTimeout(total=1800, connect=60, sock_read=1800, sock_connect=60)) as r: - assert r.status == 200, f"Failed to download {path} from {url}: {r.status}" - length = int(r.headers.get('content-length', 0)) - n_read = 0 - async with aiofiles.tempfile.NamedTemporaryFile(dir=target_dir, delete=False) as temp_file: - while chunk := await r.content.read(1024 * 1024): on_progress(n_read := n_read + await temp_file.write(chunk), length) - await aios.rename(temp_file.name, target_dir/path) - return target_dir/path + temp_file_name = None + try: + if (target_dir/path).exists(): return target_dir/path + await aios.makedirs((target_dir/path).parent, exist_ok=True) + base_url = f"{get_hf_endpoint()}/{repo_id}/resolve/{revision}/" + url = urljoin(base_url, path) + headers = await get_auth_headers() + async with aiohttp.ClientSession(timeout=aiohttp.ClientTimeout(total=1800, connect=60, sock_read=1800, sock_connect=60)) as session: + async with session.get(url, headers=headers, timeout=aiohttp.ClientTimeout(total=1800, connect=60, sock_read=1800, sock_connect=60)) as r: + assert r.status == 200, f"Failed to download {path} from {url}: {r.status}" + length = int(r.headers.get('content-length', 0)) + n_read = 0 + async with aiofiles.tempfile.NamedTemporaryFile(dir=target_dir, delete=False) as temp_file: + temp_file_name = temp_file.name + while chunk := await r.content.read(1024 * 1024): on_progress(n_read := n_read + await temp_file.write(chunk), length) + await aios.rename(temp_file.name, target_dir/path) + return target_dir/path + finally: + if temp_file_name: # attempt to delete tmp file if it still exists + try: await aios.unlink(temp_file_name) + except: pass + def calculate_repo_progress(shard: Shard, repo_id: str, revision: str, file_progress: Dict[str, RepoFileProgressEvent], all_start_time: float) -> RepoProgressEvent: all_total_bytes = sum([p.total for p in file_progress.values()]) @@ -233,4 +241,3 @@ class NewShardDownloader(ShardDownloader): if DEBUG >= 6: print("Downloaded shards:", downloads) if any(isinstance(d, Exception) for d in downloads) and DEBUG >= 1: print("Error downloading shards:", [d for d in downloads if isinstance(d, Exception)]) return [d for d in downloads if not isinstance(d, Exception)] - From 0bebf8dfdea9ad24c9437de46f006d3dfd6e6815 Mon Sep 17 00:00:00 2001 From: Alex Cheema Date: Thu, 30 Jan 2025 20:21:28 +0000 Subject: [PATCH 2/2] fix indent --- exo/download/new_shard_download.py | 26 +++++++++++++------------- 1 file changed, 13 insertions(+), 13 deletions(-) diff --git a/exo/download/new_shard_download.py b/exo/download/new_shard_download.py index e4efb3df..0e8957ae 100644 --- a/exo/download/new_shard_download.py +++ b/exo/download/new_shard_download.py @@ -75,7 +75,7 @@ async def fetch_file_list(repo_id, revision, path=""): url = f"{api_url}/{path}" if path else api_url headers = await get_auth_headers() - async with aiohttp.ClientSession(timeout=aiohttp.ClientTimeout(total=30, connect=10, sock_read=1800, sock_connect=60)) as session: + async with aiohttp.ClientSession(timeout=aiohttp.ClientTimeout(total=30, connect=10, sock_read=30, sock_connect=10)) as session: async with session.get(url, headers=headers) as response: if response.status == 200: data = await response.json() @@ -84,7 +84,7 @@ async def fetch_file_list(repo_id, revision, path=""): if item["type"] == "file": files.append({"path": item["path"], "size": item["size"]}) elif item["type"] == "directory": - subfiles = await fetch_file_list(session, repo_id, revision, item["path"]) + subfiles = await fetch_file_list(repo_id, revision, item["path"]) files.extend(subfiles) return files else: @@ -169,17 +169,17 @@ async def download_shard(shard: Shard, inference_engine_classname: str, on_progr downloaded_bytes = (await aios.stat(target_dir/file["path"])).st_size if await aios.path.exists(target_dir/file["path"]) else 0 file_progress[file["path"]] = RepoFileProgressEvent(repo_id, revision, file["path"], downloaded_bytes, 0, file["size"], 0, timedelta(0), "complete" if downloaded_bytes == file["size"] else "not_started", time.time()) - semaphore = asyncio.Semaphore(max_parallel_downloads) - async def download_with_semaphore(file): - async with semaphore: - await download_file(repo_id, revision, file["path"], target_dir, lambda curr_bytes, total_bytes: on_progress_wrapper(file, curr_bytes, total_bytes)) - if not skip_download: await asyncio.gather(*[download_with_semaphore(file) for file in filtered_file_list]) - final_repo_progress = calculate_repo_progress(shard, repo_id, revision, file_progress, all_start_time) - on_progress.trigger_all(shard, final_repo_progress) - if gguf := next((f for f in filtered_file_list if f["path"].endswith(".gguf")), None): - return target_dir/gguf["path"], final_repo_progress - else: - return target_dir, final_repo_progress + semaphore = asyncio.Semaphore(max_parallel_downloads) + async def download_with_semaphore(file): + async with semaphore: + await download_file(repo_id, revision, file["path"], target_dir, lambda curr_bytes, total_bytes: on_progress_wrapper(file, curr_bytes, total_bytes)) + if not skip_download: await asyncio.gather(*[download_with_semaphore(file) for file in filtered_file_list]) + final_repo_progress = calculate_repo_progress(shard, repo_id, revision, file_progress, all_start_time) + on_progress.trigger_all(shard, final_repo_progress) + if gguf := next((f for f in filtered_file_list if f["path"].endswith(".gguf")), None): + return target_dir/gguf["path"], final_repo_progress + else: + return target_dir, final_repo_progress def new_shard_downloader() -> ShardDownloader: return SingletonShardDownloader(CachedShardDownloader(NewShardDownloader()))