feat: learn framework architecture incompatibilities
This commit is contained in:
65
modelhub_submmit_api/architecture_compatibility.py
Normal file
65
modelhub_submmit_api/architecture_compatibility.py
Normal file
@@ -0,0 +1,65 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Iterable
|
||||
from typing import Any
|
||||
|
||||
|
||||
EXPLICIT_ARCHITECTURE_FAILURE_CATEGORY = "framework_architecture_unsupported"
|
||||
EXPLICIT_ARCHITECTURE_FAILURE_ACTION = "block_gpu_framework_architecture"
|
||||
EXPLICIT_ARCHITECTURE_FAILURE_REASON = "explicit_framework_model_unsupported"
|
||||
DEFAULT_ARCHITECTURE_BLOCK_TTL_DAYS = 30
|
||||
|
||||
|
||||
def architecture_profile(
|
||||
model_type: Any,
|
||||
architectures: Any,
|
||||
) -> dict[str, Any] | None:
|
||||
"""Build a stable, conservative architecture identity for feedback matching."""
|
||||
normalized_architectures = _normalize_architectures(architectures)
|
||||
normalized_model_type = _normalize(model_type)
|
||||
if normalized_architectures:
|
||||
return {
|
||||
"matchType": "architectures",
|
||||
"signature": "architectures:" + ",".join(normalized_architectures),
|
||||
"architectures": normalized_architectures,
|
||||
"modelType": normalized_model_type or None,
|
||||
}
|
||||
if normalized_model_type:
|
||||
return {
|
||||
"matchType": "model_type",
|
||||
"signature": f"model_type:{normalized_model_type}",
|
||||
"architectures": [],
|
||||
"modelType": normalized_model_type,
|
||||
}
|
||||
return None
|
||||
|
||||
|
||||
def architecture_compatibility_key(
|
||||
target_gpu: Any,
|
||||
framework: Any,
|
||||
task_type: Any,
|
||||
signature: Any,
|
||||
) -> str | None:
|
||||
parts = (
|
||||
_normalize(target_gpu),
|
||||
_normalize(framework),
|
||||
_normalize(task_type),
|
||||
_normalize(signature),
|
||||
)
|
||||
if not all(parts):
|
||||
return None
|
||||
return "|".join(parts)
|
||||
|
||||
|
||||
def _normalize_architectures(value: Any) -> list[str]:
|
||||
if isinstance(value, str):
|
||||
values: Iterable[Any] = [value]
|
||||
elif isinstance(value, (list, tuple, set)):
|
||||
values = value
|
||||
else:
|
||||
values = []
|
||||
return sorted({_normalize(item) for item in values if _normalize(item)})
|
||||
|
||||
|
||||
def _normalize(value: Any) -> str:
|
||||
return str(value or "").strip().casefold()
|
||||
Reference in New Issue
Block a user