ActionCodec2-2nd-order / processing_actioncodec2.py
ZibinDong's picture
Upload pretrained ActionCodec2 artifact
fee0e43 verified
Raw History Blame Contribute Delete
3.41 kB
"""Small Hugging Face entry point for a grouped ActionCodec2 artifact.
Transformers copies only top-level Python files from a Hub repository into its
dynamic-module cache. This entry point locates the artifact, then loads the
versioned runtime stored in ``runtime/`` without installing another package.
"""
from __future__ import annotations
import hashlib
import importlib
import importlib.util
import sys
from pathlib import Path
class ActionCodec2:
"""Load the processor implementation bundled with a pretrained artifact."""
@classmethod
def register_for_auto_class(cls, auto_class="AutoProcessor"):
"""Satisfy Transformers' dynamic class hook for this thin loader."""
return cls
@classmethod
def from_pretrained(cls, pretrained_model_name_or_path, **kwargs):
source = Path(pretrained_model_name_or_path)
subfolder = str(kwargs.pop("subfolder", ""))
action_space = kwargs.pop("action_space", None)
kwargs.pop("_from_auto", None)
kwargs.pop("trust_remote_code", None)
allowed = {
"cache_dir",
"force_download",
"local_files_only",
"token",
"revision",
"repo_type",
}
unknown = sorted(set(kwargs) - allowed)
if unknown:
raise TypeError(
f"unsupported from_pretrained keyword(s): {', '.join(unknown)}"
)
if source.is_dir():
root = source / subfolder
else:
from huggingface_hub import snapshot_download
prefix = f"{subfolder.rstrip('/')}/" if subfolder else ""
patterns = [
"config.json",
"processor_config.json",
"router_config.yaml",
"profiles/**",
"runtime/**",
"README.md",
"requirements.txt",
"fit_report.json",
]
snapshot = snapshot_download(
repo_id=str(pretrained_model_name_or_path),
allow_patterns=[prefix + pattern for pattern in patterns],
**kwargs,
)
root = Path(snapshot) / subfolder
runtime = root / "runtime"
package_file = runtime / "__init__.py"
if not package_file.is_file():
raise FileNotFoundError(
f"ActionCodec2 runtime is missing from {root}; copy the whole artifact"
)
identity = hashlib.sha256(str(runtime.resolve()).encode()).hexdigest()[:16]
package_name = f"_actioncodec2_artifact_{identity}"
if package_name not in sys.modules:
spec = importlib.util.spec_from_file_location(
package_name, package_file, submodule_search_locations=[str(runtime)]
)
if spec is None or spec.loader is None:
raise ImportError(f"cannot load ActionCodec2 runtime from {runtime}")
package = importlib.util.module_from_spec(spec)
sys.modules[package_name] = package
try:
spec.loader.exec_module(package)
except BaseException:
del sys.modules[package_name]
raise
processor = importlib.import_module(f"{package_name}.processing_actioncodec2")
return processor.ActionCodec2.from_pretrained(root, action_space=action_space)