← back to Exo
Refactor model types
cb101e3d24c28ee5120db36bd01142c468f92983 · 2025-07-21 20:35:27 +0100 · Seth Howes
Files touched
M master/api.pyM master/main.pyM shared/types/events/chunks.pyR069 shared/types/models/metadata.py shared/types/models.pyD shared/types/models/common.pyD shared/types/models/model.pyD shared/types/models/sources.pyM shared/types/worker/downloads.pyM shared/types/worker/runners.pyM shared/types/worker/shards.pyM worker/tests/conftest.pyM worker/tests/test_worker_plan.py
Diff
commit cb101e3d24c28ee5120db36bd01142c468f92983
Author: Seth Howes <71157822+sethhowes@users.noreply.github.com>
Date: Mon Jul 21 20:35:27 2025 +0100
Refactor model types
---
master/api.py | 8 ++--
master/main.py | 5 +-
shared/types/events/chunks.py | 2 +-
shared/types/{models/metadata.py => models.py} | 6 ++-
shared/types/models/common.py | 5 --
shared/types/models/model.py | 17 -------
shared/types/models/sources.py | 66 --------------------------
shared/types/worker/downloads.py | 4 +-
shared/types/worker/runners.py | 2 +-
shared/types/worker/shards.py | 2 +-
worker/tests/conftest.py | 4 +-
worker/tests/test_worker_plan.py | 2 +-
12 files changed, 16 insertions(+), 107 deletions(-)
diff --git a/master/api.py b/master/api.py
index 0bbc2fbd..2751f2df 100644
--- a/master/api.py
+++ b/master/api.py
@@ -1,9 +1,7 @@
from typing import Protocol
from shared.types.graphs.topology import Topology
-from shared.types.models.common import ModelId
-from shared.types.models.model import ModelInfo
-from shared.types.models.sources import ModelSource
+from shared.types.models import ModelId, ModelMetadata
from shared.types.worker.common import InstanceId
from shared.types.worker.downloads import DownloadProgress
from shared.types.worker.instances import Instance
@@ -20,8 +18,8 @@ class ClusterAPI(Protocol):
def remove_instance(self, instance_id: InstanceId) -> None: ...
- def get_model_data(self, model_id: ModelId) -> ModelInfo: ...
+ def get_model_metadata(self, model_id: ModelId) -> ModelMetadata: ...
- def download_model(self, model_id: ModelId, model_source: ModelSource) -> None: ...
+ def download_model(self, model_id: ModelId) -> None: ...
def get_download_progress(self, model_id: ModelId) -> DownloadProgress: ...
diff --git a/master/main.py b/master/main.py
index a8fd53ca..9a131e0e 100644
--- a/master/main.py
+++ b/master/main.py
@@ -22,8 +22,7 @@ from shared.logger import (
from shared.types.events.common import (
EventCategoryEnum,
)
-from shared.types.models.common import ModelId
-from shared.types.models.model import ModelInfo
+from shared.types.models import ModelId, ModelMetadata
from shared.types.state import State
from shared.types.worker.common import InstanceId
from shared.types.worker.instances import Instance
@@ -180,7 +179,7 @@ def remove_instance(instance_id: InstanceId) -> None: ...
@app.get("/model/{model_id}/metadata")
-def get_model_data(model_id: ModelId) -> ModelInfo: ...
+def get_model_metadata(model_id: ModelId) -> ModelMetadata: ...
@app.post("/model/{model_id}/instances")
diff --git a/shared/types/events/chunks.py b/shared/types/events/chunks.py
index 65bf4dd6..8db92f51 100644
--- a/shared/types/events/chunks.py
+++ b/shared/types/events/chunks.py
@@ -6,7 +6,7 @@ from typing import Annotated, Literal
from pydantic import BaseModel, Field, TypeAdapter
from shared.openai_compat import FinishReason
-from shared.types.models.common import ModelId
+from shared.types.models import ModelId
from shared.types.tasks.common import TaskId
diff --git a/shared/types/models/metadata.py b/shared/types/models.py
similarity index 69%
rename from shared/types/models/metadata.py
rename to shared/types/models.py
index 1c0015e9..3d3d0456 100644
--- a/shared/types/models/metadata.py
+++ b/shared/types/models.py
@@ -1,10 +1,12 @@
-from typing import Annotated, final
+from typing import Annotated, TypeAlias
from pydantic import BaseModel, PositiveInt
+ModelId: TypeAlias = str
+
-@final
class ModelMetadata(BaseModel):
+ model_id: ModelId
pretty_name: str
storage_size_kilobytes: Annotated[int, PositiveInt]
n_layers: Annotated[int, PositiveInt]
diff --git a/shared/types/models/common.py b/shared/types/models/common.py
deleted file mode 100644
index 05e82a34..00000000
--- a/shared/types/models/common.py
+++ /dev/null
@@ -1,5 +0,0 @@
-from shared.types.common import NewUUID
-
-
-class ModelId(NewUUID):
- pass
diff --git a/shared/types/models/model.py b/shared/types/models/model.py
deleted file mode 100644
index c50ade27..00000000
--- a/shared/types/models/model.py
+++ /dev/null
@@ -1,17 +0,0 @@
-from typing import Sequence, final
-
-from pydantic import BaseModel, TypeAdapter
-
-from shared.types.models.common import ModelId
-from shared.types.models.metadata import ModelMetadata
-from shared.types.models.sources import ModelSource
-
-
-@final
-class ModelInfo(BaseModel):
- model_id: ModelId
- model_sources: Sequence[ModelSource]
- model_metadata: ModelMetadata
-
-
-ModelIdAdapter: TypeAdapter[ModelId] = TypeAdapter(ModelId)
diff --git a/shared/types/models/sources.py b/shared/types/models/sources.py
deleted file mode 100644
index a3712bff..00000000
--- a/shared/types/models/sources.py
+++ /dev/null
@@ -1,66 +0,0 @@
-from enum import Enum
-from typing import Annotated, Any, Literal, Union, final
-
-from pydantic import AnyHttpUrl, BaseModel, Field, TypeAdapter
-
-from shared.types.models.common import ModelId
-
-
-@final
-class SourceType(str, Enum):
- HuggingFace = "HuggingFace"
- GitHub = "GitHub"
-
-
-@final
-class SourceFormatType(str, Enum):
- HuggingFaceTransformers = "HuggingFaceTransformers"
-
-
-RepoPath = Annotated[str, Field(pattern=r"^[^/]+/[^/]+$")]
-
-
-class BaseModelSource[T: SourceType, S: SourceFormatType](BaseModel):
- model_uuid: ModelId
- source_type: T
- source_format: S
- source_data: Any
-
-
-@final
-class HuggingFaceModelSourceData(BaseModel):
- path: RepoPath
-
-
-@final
-class GitHubModelSourceData(BaseModel):
- url: AnyHttpUrl
-
-
-@final
-class HuggingFaceModelSource(
- BaseModelSource[SourceType.HuggingFace, SourceFormatType.HuggingFaceTransformers]
-):
- source_type: Literal[SourceType.HuggingFace] = SourceType.HuggingFace
- source_format: Literal[SourceFormatType.HuggingFaceTransformers] = (
- SourceFormatType.HuggingFaceTransformers
- )
- source_data: HuggingFaceModelSourceData
-
-
-@final
-class GitHubModelSource(BaseModelSource[SourceType.GitHub, SourceFormatType]):
- source_type: Literal[SourceType.GitHub] = SourceType.GitHub
- source_format: SourceFormatType
- source_data: GitHubModelSourceData
-
-
-_ModelSource = Annotated[
- Union[
- HuggingFaceModelSource,
- GitHubModelSource,
- ],
- Field(discriminator="source_type"),
-]
-ModelSource = BaseModelSource[SourceType, SourceFormatType]
-ModelSourceAdapter: TypeAdapter[ModelSource] = TypeAdapter(_ModelSource)
diff --git a/shared/types/worker/downloads.py b/shared/types/worker/downloads.py
index 649eb48b..a9e40c19 100644
--- a/shared/types/worker/downloads.py
+++ b/shared/types/worker/downloads.py
@@ -11,8 +11,7 @@ from typing import (
from pydantic import BaseModel, Field, PositiveInt
from shared.types.common import NodeId
-from shared.types.models.common import ModelId
-from shared.types.models.sources import ModelSource
+from shared.types.models import ModelId
from shared.types.worker.shards import ShardMetadata
@@ -74,7 +73,6 @@ DownloadEffectHandler = Callable[
def download_shard(
model_id: ModelId,
- model_source: ModelSource,
shard_metadata: ShardMetadata,
effect_handlers: Sequence[DownloadEffectHandler],
) -> None: ...
diff --git a/shared/types/worker/runners.py b/shared/types/worker/runners.py
index 1b6c371b..51a08958 100644
--- a/shared/types/worker/runners.py
+++ b/shared/types/worker/runners.py
@@ -5,7 +5,7 @@ from typing import Annotated, Generic, Literal, TypeVar
from pydantic import BaseModel, Field, TypeAdapter, model_validator
from shared.types.common import NodeId
-from shared.types.models.common import ModelId
+from shared.types.models import ModelId
from shared.types.worker.common import RunnerId
from shared.types.worker.downloads import DownloadProgress
from shared.types.worker.shards import ShardMetadata
diff --git a/shared/types/worker/shards.py b/shared/types/worker/shards.py
index 5ee7baa8..3decee54 100644
--- a/shared/types/worker/shards.py
+++ b/shared/types/worker/shards.py
@@ -4,7 +4,7 @@ from typing import Annotated, Generic, Literal, TypeAlias, TypeVar
from pydantic import BaseModel, DirectoryPath, Field, TypeAdapter
from shared.types.common import NodeId
-from shared.types.models.common import ModelId
+from shared.types.models import ModelId
class PartitionStrategy(str, Enum):
diff --git a/worker/tests/conftest.py b/worker/tests/conftest.py
index afe312c0..4fae4868 100644
--- a/worker/tests/conftest.py
+++ b/worker/tests/conftest.py
@@ -7,7 +7,7 @@ from typing import Callable
import pytest
from shared.types.common import NodeId
-from shared.types.models.common import ModelId
+from shared.types.models import ModelId
from shared.types.state import State
from shared.types.tasks.common import (
ChatCompletionMessage,
@@ -45,7 +45,7 @@ def pipeline_shard_meta():
return PipelineShardMetadata(
device_rank=device_rank,
- model_id=ModelId(uuid=uuid.uuid4()),
+ model_id=ModelId(uuid.uuid4()),
model_path=Path(
"~/.exo/models/mlx-community--Llama-3.2-1B-Instruct-4bit/"
).expanduser(),
diff --git a/worker/tests/test_worker_plan.py b/worker/tests/test_worker_plan.py
index 56c0503b..cdc59623 100644
--- a/worker/tests/test_worker_plan.py
+++ b/worker/tests/test_worker_plan.py
@@ -8,7 +8,7 @@ from typing import Final, List, Optional, Type
import pytest
from shared.types.common import NodeId
-from shared.types.models.common import ModelId
+from shared.types.models import ModelId
from shared.types.state import State
# WorkerState import below after RunnerCase definition to avoid forward reference issues
← 54efd01d add forwarder supervisor
·
back to Exo
·
Downloads 449fdac2 →