harrrshall's picture
Release BarunAction-35M candidate-v2
5a46e5d verified
Raw
History Blame Contribute Delete
3.83 kB
"""Local proposal-only inference package for BarunAction-35M.
Public objects are loaded on first access so lightweight operations such as package
discovery and checkpoint download do not import the model runtime unnecessarily.
"""
from __future__ import annotations
from importlib import import_module
from typing import Any
_CANDIDATE_EXPORTS = (
"BASE_MODEL_NAME",
"CANDIDATE_ARM_ID",
"CANDIDATE_CHECKPOINT_SHA256",
"CANDIDATE_ID",
"CANDIDATE_MANIFEST_SHA256",
"CANDIDATE_REPO_ID",
"CANDIDATE_REVISION",
"CANDIDATE_RUN_ID",
"CANDIDATE_SELECTION_RUN_ID",
"CANDIDATE_STEP",
"MODEL_NAME",
"candidate_identity",
)
_HUB_EXPORTS = (
"CHECKPOINT_MANIFEST_NAME",
"DOWNLOAD_ALLOW_PATTERNS",
"DOWNLOAD_SCHEMA_VERSION",
"PUBLIC_MODEL_ID",
"PUBLIC_MODEL_REVISION",
"DownloadedCheckpoint",
"HubDownloadError",
"download_candidate_checkpoint",
)
_INFERENCE_EXPORTS = (
"DEFAULT_MAX_NEW_TOKENS",
"PROMPT_CONTRACT_VERSION",
"RESULT_SCHEMA_VERSION",
"BarunActionCompiler",
"ErrorDetail",
"InferenceOutcome",
"PolicyAssessment",
"assess_policy",
"render_prompt",
"validate_action_output",
)
_QUANTIZATION_EXPORTS = (
"INT8_SMOKE_CASES_VERSION",
"INT8_SMOKE_REPORT_VERSION",
"ActionIRSmokeCase",
"QuantizationSmokeError",
"compare_int8_action_ir",
"parse_int8_smoke_cases",
)
_SCHEMA_EXPORTS = (
"SCHEMA_CONTRACT_VERSION",
"ContractError",
"PreparedInput",
"ToolDeclaration",
"canonical_context",
"canonical_now",
"parse_tool_declarations",
"prepare_input",
)
_SIMULATOR_EXPORTS = (
"SIMULATOR_SCHEMA_VERSION",
"SimulationResult",
"simulate_action",
)
_LAZY_EXPORTS = {
**{name: (".candidate", name) for name in _CANDIDATE_EXPORTS},
**{name: (".hub", name) for name in _HUB_EXPORTS},
**{name: (".inference", name) for name in _INFERENCE_EXPORTS},
**{name: (".quantization", name) for name in _QUANTIZATION_EXPORTS},
**{name: (".schema", name) for name in _SCHEMA_EXPORTS},
**{name: (".simulator", name) for name in _SIMULATOR_EXPORTS},
}
__all__ = [
"BASE_MODEL_NAME",
"CANDIDATE_ARM_ID",
"CANDIDATE_CHECKPOINT_SHA256",
"CANDIDATE_ID",
"CANDIDATE_MANIFEST_SHA256",
"CANDIDATE_REPO_ID",
"CANDIDATE_REVISION",
"CANDIDATE_RUN_ID",
"CANDIDATE_SELECTION_RUN_ID",
"CANDIDATE_STEP",
"CHECKPOINT_MANIFEST_NAME",
"DEFAULT_MAX_NEW_TOKENS",
"DOWNLOAD_ALLOW_PATTERNS",
"DOWNLOAD_SCHEMA_VERSION",
"INT8_SMOKE_CASES_VERSION",
"INT8_SMOKE_REPORT_VERSION",
"MODEL_NAME",
"PROMPT_CONTRACT_VERSION",
"PUBLIC_MODEL_ID",
"PUBLIC_MODEL_REVISION",
"RESULT_SCHEMA_VERSION",
"SCHEMA_CONTRACT_VERSION",
"SIMULATOR_SCHEMA_VERSION",
"ActionIRSmokeCase",
"BarunActionCompiler",
"ContractError",
"DownloadedCheckpoint",
"ErrorDetail",
"HubDownloadError",
"InferenceOutcome",
"PolicyAssessment",
"PreparedInput",
"QuantizationSmokeError",
"SimulationResult",
"ToolDeclaration",
"assess_policy",
"candidate_identity",
"canonical_context",
"canonical_now",
"compare_int8_action_ir",
"download_candidate_checkpoint",
"parse_int8_smoke_cases",
"parse_tool_declarations",
"prepare_input",
"render_prompt",
"simulate_action",
"validate_action_output",
]
def __getattr__(name: str) -> Any:
try:
module_name, attribute_name = _LAZY_EXPORTS[name]
except KeyError as error:
raise AttributeError(f"module {__name__!r} has no attribute {name!r}") from error
value = getattr(import_module(module_name, __name__), attribute_name)
globals()[name] = value
return value
def __dir__() -> list[str]:
return sorted(set(globals()) | set(__all__))