[object Object]

← back to Exo

Refactor model types

cb101e3d24c28ee5120db36bd01142c468f92983 · 2025-07-21 20:35:27 +0100 · Seth Howes

Files touched

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 →