mac-compute-space / app /capabilities.py
josephrw's picture
Upload folder using huggingface_hub
c4916b2 verified
Raw
History Blame Contribute Delete
2.4 kB
from typing import Dict, List, Optional
from app.models import DeviceCapability, WorkerState, JobType
_capability_registry: Dict[str, List[Dict[str, any]]] = {}
CAPABILITY_NAMES = {
JobType.TEXT_EMBEDDING: "iphone.text.embedding.private",
JobType.IMAGE_CLASSIFICATION: "iphone.image.classify.local",
JobType.IMAGE_EMBEDDING: "iphone.image.embed.local",
JobType.LOCAL_OCR: "iphone.ocr.local",
JobType.AUDIO_TRANSCRIPTION: "iphone.audio.transcribe.local",
JobType.SMALL_LLM_GENERATE: "iphone.llm.generate.local",
JobType.PRIVACY_REDACTION: "iphone.privacy.redact.local",
JobType.SENSOR_CLASSIFICATION: "iphone.sensor.motion.classify",
}
def normalize_capability(raw: dict) -> DeviceCapability:
return DeviceCapability(
capability_name=raw.get("capability_name", ""),
runtime_type=raw.get("runtime_type", "safari_wasm"),
model_id=raw.get("model_id"),
model_hash=raw.get("model_hash"),
quantization=raw.get("quantization"),
max_input_bytes=raw.get("max_input_bytes"),
estimated_latency_ms=raw.get("estimated_latency_ms"),
)
def validate_capability(capability: DeviceCapability) -> bool:
return capability.capability_name in CAPABILITY_NAMES.values()
def register_capabilities(session_id: str, worker_id: str, capabilities: List[dict]) -> bool:
key = f"{session_id}:{worker_id}"
_capability_registry[key] = capabilities
return True
def get_capabilities(session_id: str) -> List[Dict[str, any]]:
result = []
for key, caps in _capability_registry.items():
if key.startswith(f"{session_id}:"):
for c in caps:
result.append({**c, "source_worker": key.split(":", 1)[1]})
return result
def find_workers_with_capability(session_id: str, job_type: JobType) -> List[str]:
target = CAPABILITY_NAMES.get(job_type)
workers = []
for key, caps in _capability_registry.items():
if key.startswith(f"{session_id}:"):
for c in caps:
if c.get("capability_name") == target:
workers.append(key.split(":", 1)[1])
break
return workers
def capability_to_endpoint_path(session_id: str, capability_name: str) -> str:
return f"/api/session/{session_id}/cap/{capability_name}"
def get_capability_names() -> List[str]:
return list(CAPABILITY_NAMES.values())