77 lines
3.4 KiB
Python
77 lines
3.4 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 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()
|