airfRANS-model-exploration/tests/test_hf_upload.py

78 lines
3.4 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 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]] = []
uploaded: list[tuple[str, str, str]] = []
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 upload_file(self, *, repo_id: str, repo_type: str, path_or_fileobj: str, path_in_repo: str, commit_message: str):
uploaded.append((repo_id, repo_type, path_in_repo))
return types.SimpleNamespace(commit_url="https://huggingface.co/repo/commit/abc", oid="abc")
fake_module = types.SimpleNamespace(HfApi=FakeApi)
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(uploaded, [("owner/repo", "model", "runs/model/run-1/checkpoint_latest.pt")])
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_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()