airfRANS-model-exploration/tests/test_hf_upload.py

121 lines
5.7 KiB
Python
Raw Normal View History

2026-07-25 16:12:49 +00:00
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
2026-07-25 16:12:49 +00:00
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
2026-07-25 16:12:49 +00:00
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))
2026-07-25 16:12:49 +00:00
return types.SimpleNamespace(commit_url="https://huggingface.co/repo/commit/abc", oid="abc")
fake_module = types.SimpleNamespace(HfApi=FakeApi, CommitOperationAdd=FakeCommitOperationAdd)
2026-07-25 16:12:49 +00:00
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")])
2026-07-25 16:12:49 +00:00
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)
2026-07-25 16:12:49 +00:00
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()