120 lines
5.7 KiB
Python
120 lines
5.7 KiB
Python
from __future__ import annotations
|
|
|
|
import json
|
|
import os
|
|
import sys
|
|
import tempfile
|
|
import types
|
|
import unittest
|
|
from pathlib import Path
|
|
from unittest.mock import Mock, patch
|
|
|
|
from airfrans_frontier.training.hf_upload import HfArtifactUploader, resolve_resume_checkpoint
|
|
|
|
|
|
class HuggingFaceUploadTests(unittest.TestCase):
|
|
def test_uploader_creates_repo_uploads_file_and_writes_manifest(self) -> None:
|
|
created: list[tuple[str, str, bool]] = []
|
|
committed: list[tuple[str, str, tuple[str, ...], str]] = []
|
|
|
|
class FakeCommitOperationAdd:
|
|
def __init__(self, *, path_in_repo: str, path_or_fileobj: str) -> None:
|
|
self.path_in_repo = path_in_repo
|
|
self.path_or_fileobj = path_or_fileobj
|
|
|
|
class FakeApi:
|
|
def __init__(self, token: str) -> None:
|
|
self.token = token
|
|
|
|
def create_repo(self, *, repo_id: str, repo_type: str, private: bool, exist_ok: bool) -> None:
|
|
created.append((repo_id, repo_type, private))
|
|
|
|
def create_commit(self, *, repo_id: str, repo_type: str, operations: list[FakeCommitOperationAdd], commit_message: str):
|
|
committed.append((repo_id, repo_type, tuple(operation.path_in_repo for operation in operations), commit_message))
|
|
return types.SimpleNamespace(commit_url="https://huggingface.co/repo/commit/abc", oid="abc")
|
|
|
|
fake_module = types.SimpleNamespace(HfApi=FakeApi, CommitOperationAdd=FakeCommitOperationAdd)
|
|
with tempfile.TemporaryDirectory() as tmp, patch.dict(sys.modules, {"huggingface_hub": fake_module}), patch.dict(os.environ, {"HF_TOKEN": "token"}):
|
|
root = Path(tmp)
|
|
(root / "checkpoint_latest.pt").write_bytes(b"checkpoint")
|
|
uploader = HfArtifactUploader(
|
|
enabled=True,
|
|
run_dir=root,
|
|
repo_id="owner/repo",
|
|
repo_type="model",
|
|
path_in_repo="runs/model/run-1",
|
|
private=False,
|
|
)
|
|
|
|
result = uploader.upload_files(("checkpoint_latest.pt",), commit_message="upload checkpoint")
|
|
|
|
self.assertEqual(result["uploaded"], ["runs/model/run-1/checkpoint_latest.pt"])
|
|
self.assertEqual(created, [("owner/repo", "model", False)])
|
|
self.assertEqual(committed, [("owner/repo", "model", ("runs/model/run-1/checkpoint_latest.pt",), "upload checkpoint")])
|
|
manifest = json.loads((root / "hf_upload_manifest.json").read_text())
|
|
self.assertTrue(manifest["enabled"])
|
|
self.assertEqual(manifest["repo_id"], "owner/repo")
|
|
self.assertIn("runs/model/run-1/checkpoint_latest.pt", manifest["uploaded_paths"])
|
|
|
|
def test_uploader_suppresses_uploads_after_hf_retry_after_limit(self) -> None:
|
|
class FakeRateLimitError(RuntimeError):
|
|
def __init__(self) -> None:
|
|
super().__init__("429 Too Many Requests: Retry after 600 seconds")
|
|
self.response = types.SimpleNamespace(headers={"Retry-After": "600"})
|
|
|
|
class FakeCommitOperationAdd:
|
|
def __init__(self, *, path_in_repo: str, path_or_fileobj: str) -> None:
|
|
self.path_in_repo = path_in_repo
|
|
self.path_or_fileobj = path_or_fileobj
|
|
|
|
fake_module = types.SimpleNamespace(CommitOperationAdd=FakeCommitOperationAdd)
|
|
with tempfile.TemporaryDirectory() as tmp, patch.dict(sys.modules, {"huggingface_hub": fake_module}):
|
|
root = Path(tmp)
|
|
(root / "metrics.jsonl").write_text("{}\n")
|
|
uploader = HfArtifactUploader(
|
|
enabled=True,
|
|
run_dir=root,
|
|
repo_id="owner/repo",
|
|
repo_type="model",
|
|
path_in_repo="runs/model/run-1",
|
|
private=False,
|
|
max_rate_limit_sleep_seconds=0,
|
|
)
|
|
fake_api = types.SimpleNamespace(create_commit=Mock(side_effect=FakeRateLimitError()))
|
|
uploader._api = fake_api
|
|
|
|
with self.assertRaises(FakeRateLimitError):
|
|
uploader.upload_files(("metrics.jsonl",), commit_message="first")
|
|
suppressed = uploader.upload_files(("metrics.jsonl",), commit_message="second")
|
|
|
|
self.assertTrue(suppressed["rate_limited"])
|
|
self.assertEqual(fake_api.create_commit.call_count, 1)
|
|
manifest = json.loads((root / "hf_upload_manifest.json").read_text())
|
|
self.assertGreater(manifest["rate_limit_until"], 0)
|
|
self.assertEqual(manifest["rate_limit_retry_after_seconds"], 600.0)
|
|
self.assertEqual(len(manifest["suppressed_uploads"]), 2)
|
|
|
|
def test_resolve_resume_checkpoint_downloads_hf_uri(self) -> None:
|
|
calls: list[tuple[str, str]] = []
|
|
|
|
def fake_download(*, repo_id: str, repo_type: str, filename: str, token: str, local_dir: str) -> str:
|
|
calls.append((repo_id, filename))
|
|
path = Path(local_dir) / filename
|
|
path.parent.mkdir(parents=True, exist_ok=True)
|
|
path.write_bytes(b"checkpoint")
|
|
return str(path)
|
|
|
|
fake_module = types.SimpleNamespace(hf_hub_download=fake_download)
|
|
with patch.dict(sys.modules, {"huggingface_hub": fake_module}), patch.dict(os.environ, {"HF_TOKEN": "token"}):
|
|
path, info = resolve_resume_checkpoint("hf://owner/repo/runs/model/checkpoint_latest.pt")
|
|
|
|
self.assertIsNotNone(path)
|
|
assert path is not None
|
|
self.assertTrue(path.is_file())
|
|
self.assertEqual(calls, [("owner/repo", "runs/model/checkpoint_latest.pt")])
|
|
self.assertTrue(info["resume_downloaded"])
|
|
self.assertEqual(info["resume_source"], "hf://owner/repo/runs/model/checkpoint_latest.pt")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|