diff --git a/.github/workflows/package.yml b/.github/workflows/package.yml new file mode 100644 index 0000000..fb94351 --- /dev/null +++ b/.github/workflows/package.yml @@ -0,0 +1,71 @@ +name: Package +on: + pull_request: + push: + branches: [main] + workflow_call: + inputs: + release: + type: boolean + default: false +permissions: + contents: read +env: + UV_DEFAULT_INDEX: https://pypi.org/simple +jobs: + build: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + - uses: actions/setup-python@v5 + with: + python-version: '3.12' + - run: python -m pip install uv==0.9.26 + # uv builds the wheel from the sdist, checking that source releases rebuild. + - run: uv build + - run: uvx --from twine twine check --strict dist/* + - run: uv run --no-project --with packaging python scripts/check_distribution.py + - name: Check PyPI release prerequisites + if: inputs.release + env: + RELEASE_TAG: ${{ github.event.release.tag_name }} + run: uv run --no-project --with packaging python scripts/check_distribution.py --for-pypi --tag "$RELEASE_TAG" + - uses: actions/upload-artifact@v4 + with: + name: distributions + path: dist/* + if-no-files-found: error + install: + needs: build + runs-on: ubuntu-latest + strategy: + matrix: + python: ['3.11', '3.12'] + distribution: [wheel, sdist] + steps: + - uses: actions/setup-python@v5 + with: + python-version: ${{ matrix.python }} + - uses: actions/download-artifact@v4 + with: + name: distributions + path: dist + - name: Install in a clean environment + env: + DISTRIBUTION: ${{ matrix.distribution }} + run: | + python -m venv "$RUNNER_TEMP/spindle-install" + echo "$RUNNER_TEMP/spindle-install/bin" >> "$GITHUB_PATH" + if [ "$DISTRIBUTION" = wheel ]; then + "$RUNNER_TEMP/spindle-install/bin/python" -m pip install --index-url https://pypi.org/simple dist/*.whl + else + "$RUNNER_TEMP/spindle-install/bin/python" -m pip install --index-url https://pypi.org/simple dist/*.tar.gz + fi + - name: Check installed API, CLI, and packaged presets + working-directory: ${{ runner.temp }} + run: | + python -m pip check + python -c "import spindle; from importlib.metadata import version; assert spindle.__version__ == version('modal-spindle'); from spindle.engines import qwen3_5_4b_full_64k; qwen3_5_4b_full_64k()" + spindle --help + spindle config init --preset qwen35-9b-lora-16k > deployment.py + spindle config validate deployment.py diff --git a/.github/workflows/publish.yml b/.github/workflows/publish.yml new file mode 100644 index 0000000..0c63542 --- /dev/null +++ b/.github/workflows/publish.yml @@ -0,0 +1,30 @@ +name: Publish to PyPI +on: + release: + types: [published] +permissions: + contents: read +concurrency: + group: pypi-${{ github.event.release.tag_name }} + cancel-in-progress: false +jobs: + tests: + uses: ./.github/workflows/tests.yml + package: + uses: ./.github/workflows/package.yml + with: + release: true + publish: + needs: [tests, package] + runs-on: ubuntu-latest + environment: + name: pypi + url: https://pypi.org/p/modal-spindle + permissions: + id-token: write + steps: + - uses: actions/download-artifact@v4 + with: + name: distributions + path: dist + - uses: pypa/gh-action-pypi-publish@release/v1 diff --git a/.github/workflows/tests.yml b/.github/workflows/tests.yml index d517e0b..4e510fa 100644 --- a/.github/workflows/tests.yml +++ b/.github/workflows/tests.yml @@ -1,10 +1,13 @@ name: Core CPU tests on: + workflow_call: pull_request: push: branches: [main] permissions: contents: read +env: + UV_DEFAULT_INDEX: https://pypi.org/simple jobs: test: runs-on: ubuntu-latest @@ -19,4 +22,3 @@ jobs: # so the loss, gradient, and FP32-head regressions also run in CI. - run: uv pip install --python .venv/bin/python torch==2.10.0 --index-url https://download.pytorch.org/whl/cpu - run: uv run --no-sync pytest -q - - run: uv build diff --git a/README.md b/README.md index 032ce49..5325a62 100644 --- a/README.md +++ b/README.md @@ -2,7 +2,11 @@ Spindle is a Tinker SDK-compatible backend run on Modal. Trainers run `forward_backward` and `optim_step` calls, then publish updated weights to autoscaling sampling replicas managed by the [Stitch](https://github.com/modal-projects/stitch) protocol (hence the name!). Currently, Spindle supports single-tenant full-parameter training as well as multi-tenant LoRA training. -# Getting Started +Spindle supports Python 3.11 and 3.12; use Python 3.12 for Modal deployments. +The package is being prepared for PyPI as `modal-spindle`. Until the first release, +install from Git using the instructions below. + +# Getting Started ## Full-parameter training runs @@ -23,8 +27,8 @@ with spindle.run(engine=engine) as (url, api_key): Our FFT path is *not* Tinker compatible, but roughly obeys the same abstractions. -See [scoped runs](docs/scoped-runs.md) for recovery and custom engines, -and the [Codeforces example](examples/codeforces-codegolf/README.md) for a complete +See [scoped runs](https://github.com/modal-projects/spindle/blob/main/docs/scoped-runs.md) for recovery and custom engines, +and the [Codeforces example](https://github.com/modal-projects/spindle/blob/main/examples/codeforces-codegolf/README.md) for a complete training loop with sandbox judging and checkpoints. ## LoRA training runs @@ -48,7 +52,7 @@ training = service.create_lora_training_client( ## Shared deployment quick start -Shared deployments use Python recipes inheriting from `BaseConfig`. See [Python deployment configs](docs/deployment-configs.md). Keep the active Python config list in [scripts/deploy_models.sh](scripts/deploy_models.sh); run it to deploy the complete list. +Shared deployments use Python recipes inheriting from `BaseConfig`. See [Python deployment configs](https://github.com/modal-projects/spindle/blob/main/docs/deployment-configs.md). Keep the active Python config list in [scripts/deploy_models.sh](https://github.com/modal-projects/spindle/blob/main/scripts/deploy_models.sh); run it to deploy the complete list. Install Spindle into your own Python project, deploy it once to Modal, then call its API from your training scripts. The commands below work in Bash or Zsh. @@ -137,11 +141,11 @@ uv run spindle config validate deployment.py uv run spindle deploy deployment.py ``` -This deploys the shared app and prints its `server` URL. Add more Python config files to the same command to serve more recipes. Always supply the complete current set. The Miles commit is pinned in `miles_image.py`; see [Python deployment configs](docs/deployment-configs.md). +This deploys the shared app and prints its `server` URL. Add more Python config files to the same command to serve more recipes. Always supply the complete current set. The Miles commit is pinned in `miles_image.py`; see [Python deployment configs](https://github.com/modal-projects/spindle/blob/main/docs/deployment-configs.md). From a repository checkout, maintain the list in `scripts/deploy_models.sh` and run that script. `spindle deploy` supplies the current configs and frontend platform settings to Modal. -Deploying the server doesn't allocate any GPUs; rather, this allocation for both the training and sampling sides are done on demand. See [cold starts and capacity configuration](docs/full-fine-tunes.md#performance-and-behavior-considerations) +Deploying the server doesn't allocate any GPUs; rather, this allocation for both the training and sampling sides are done on demand. See [cold starts and capacity configuration](https://github.com/modal-projects/spindle/blob/main/docs/full-fine-tunes.md#performance-and-behavior-considerations) before running a larger workload. ### 4. Clean up @@ -161,32 +165,37 @@ using `uv run modal app stop `. Stopping the frontend does not stop samp Refer to the docs for design and for more advanced features when working with either the full-parameter or LoRA paths: -Read [Working with Full Fine-Tunes](docs/full-fine-tunes.md) for full training, -or [Working with Multi-LoRA](docs/multi-lora.md) for shared Miles adapters, batch +Read [Working with Full Fine-Tunes](https://github.com/modal-projects/spindle/blob/main/docs/full-fine-tunes.md) for full training, +or [Working with Multi-LoRA](https://github.com/modal-projects/spindle/blob/main/docs/multi-lora.md) for shared Miles adapters, batch submission, scheduling, and sampling. -and the [raw Tinker RL example](scripts/rl_example.py) for sampling and a toy +and the [raw Tinker RL example](https://github.com/modal-projects/spindle/blob/main/scripts/rl_example.py) for sampling and a toy policy update. Copy examples you want to run into your project; repository `scripts/` are not installed with the package. -The [W&B RL example](scripts/wandb_rl_example.py) extends it to a multi-step +The [W&B RL example](https://github.com/modal-projects/spindle/blob/main/scripts/wandb_rl_example.py) extends it to a multi-step loop that logs reward, response length, and Spindle's training metrics to Weights & Biases from the client side; tinker-cookbook users can instead set `wandb_project`/`wandb_name` on the cookbook `Config`. -See [Design](docs/design.md) for the control-plane, training-engine, and sampling +See [Design](https://github.com/modal-projects/spindle/blob/main/docs/design.md) for the control-plane, training-engine, and sampling architecture. -See [Profiling](docs/profiling.md) for how to enable the `torch.profiler` trace of +See [Profiling](https://github.com/modal-projects/spindle/blob/main/docs/profiling.md) for how to enable the `torch.profiler` trace of a training step and read it in Perfetto. -See [Observability](docs/observability.md) for OTLP export to Datadog or a custom +See [Observability](https://github.com/modal-projects/spindle/blob/main/docs/observability.md) for OTLP export to Datadog or a custom destination, experiment labels, and the complete span/metric inventory. ## Validation -See [FFT validation](docs/validation.md) and [LoRA validation](docs/lora_validation.md) -for end-to-end training runs we've done with both parameterizations. The [Codeforces codegolf](examples/codeforces-codegolf/README.md) example provides a larger-scale e2e code-RL training run, which trains Qwen3.5-9B +See [FFT validation](https://github.com/modal-projects/spindle/blob/main/docs/validation.md) and [LoRA validation](https://github.com/modal-projects/spindle/blob/main/docs/lora_validation.md) +for end-to-end training runs we've done with both parameterizations. The [Codeforces codegolf](https://github.com/modal-projects/spindle/blob/main/examples/codeforces-codegolf/README.md) example provides a larger-scale e2e code-RL training run, which trains Qwen3.5-9B with GRPO or TailRL advantages for correctness and short solutions. It includes a sandboxed judge, checkpoint recovery, and commands to continue a checkpoint with a different reward or advantage estimator, as well as pass@k and best-of-k evaluation. + +## Development and releases + +See [Publishing](https://github.com/modal-projects/spindle/blob/main/docs/publishing.md) +for first-release prerequisites, package validation, and the PyPI release process. diff --git a/pyproject.toml b/pyproject.toml index 42cba58..dd16eb9 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,11 +1,20 @@ [build-system] -requires = ["setuptools>=68", "wheel"] +requires = ["setuptools>=77.0.3"] build-backend = "setuptools.build_meta" [project] name = "modal-spindle" version = "0.1.0" +description = "Training and disaggregated sampling on Modal with a Tinker-compatible API" +readme = "README.md" requires-python = ">=3.11,<3.13" +classifiers = [ + "Development Status :: 3 - Alpha", + "Programming Language :: Python :: 3", + "Programming Language :: Python :: 3.11", + "Programming Language :: Python :: 3.12", + "Topic :: Scientific/Engineering :: Artificial Intelligence", +] dependencies = [ "opentelemetry-exporter-otlp-proto-http>=1.39,<2", "opentelemetry-sdk>=1.39,<2", @@ -22,8 +31,15 @@ dependencies = [ "zstandard>=0.25.0", ] +[project.urls] +Homepage = "https://github.com/modal-projects/spindle" +Documentation = "https://github.com/modal-projects/spindle/blob/main/README.md" +Repository = "https://github.com/modal-projects/spindle" +Issues = "https://github.com/modal-projects/spindle/issues" + [tool.setuptools.packages.find] where = ["src"] +include = ["spindle*"] [tool.pytest.ini_options] testpaths = ["tests"] diff --git a/scripts/check_distribution.py b/scripts/check_distribution.py new file mode 100644 index 0000000..3a996cf --- /dev/null +++ b/scripts/check_distribution.py @@ -0,0 +1,91 @@ +"""Check built metadata; --for-pypi also enforces release prerequisites.""" + +import argparse +import ast +from email.parser import BytesParser +from pathlib import Path +import tarfile +import tomllib +import zipfile + +from packaging.requirements import Requirement +from packaging.specifiers import SpecifierSet + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--for-pypi", action="store_true") + parser.add_argument("--tag") + args = parser.parse_args() + root = Path(__file__).resolve().parents[1] + project = tomllib.loads((root / "pyproject.toml").read_text())["project"] + errors = [] + + version = None + for statement in ast.parse((root / "src/spindle/__init__.py").read_text()).body: + if isinstance(statement, ast.Assign) and any( + isinstance(target, ast.Name) and target.id == "__version__" + for target in statement.targets + ): + version = ast.literal_eval(statement.value) + if version != project["version"]: + errors.append("spindle.__version__ must match project.version") + if args.tag is not None and args.tag != f"v{project['version']}": + errors.append(f"Release tag must be v{project['version']}, got {args.tag!r}") + + wheels = list((root / "dist").glob("*.whl")) + sdists = list((root / "dist").glob("*.tar.gz")) + if len(wheels) != 1 or len(sdists) != 1: + raise SystemExit( + "Expected one wheel and one sdist in dist/; use a clean build directory" + ) + with zipfile.ZipFile(wheels[0]) as archive: + (metadata_path,) = ( + name for name in archive.namelist() if name.endswith(".dist-info/METADATA") + ) + wheel_metadata = BytesParser().parsebytes(archive.read(metadata_path)) + for source in (root / "src/spindle").rglob("*.py"): + if source.relative_to(root / "src").as_posix() not in archive.namelist(): + errors.append(f"Wheel is missing {source.relative_to(root)}") + with tarfile.open(sdists[0]) as archive: + (metadata_path,) = ( + member + for member in archive.getmembers() + if member.name.count("/") == 1 and member.name.endswith("/PKG-INFO") + ) + sdist_metadata = BytesParser().parsebytes( + archive.extractfile(metadata_path).read() + ) + + for kind, metadata in [("wheel", wheel_metadata), ("sdist", sdist_metadata)]: + for field, expected in [ + ("Name", "modal-spindle"), + ("Version", project["version"]), + ]: + if metadata[field] != expected: + errors.append(f"{kind}: {field} must be {expected!r}") + if SpecifierSet(metadata["Requires-Python"]) != SpecifierSet( + project["requires-python"] + ): + errors.append(f"{kind}: Requires-Python must match pyproject.toml") + if not metadata["Summary"] or not metadata.get_payload().strip(): + errors.append(f"{kind}: missing description or README") + if args.for_pypi: + if not metadata["License-Expression"] or not metadata.get_all( + "License-File" + ): + errors.append( + f"{kind}: choose a license and include its file before release" + ) + for dependency in metadata.get_all("Requires-Dist", []): + if Requirement(dependency).url: + errors.append( + f"{kind}: PyPI does not accept direct URL dependency: {dependency}" + ) + if errors: + raise SystemExit("\n".join(errors)) + print(f"Validated modal-spindle {project['version']} wheel and sdist") + + +if __name__ == "__main__": + main() diff --git a/scripts/deploy_models.sh b/scripts/deploy_models.sh index c4ca5d4..04227f3 100755 --- a/scripts/deploy_models.sh +++ b/scripts/deploy_models.sh @@ -9,9 +9,10 @@ cd "$(dirname "$0")/.." deployment_files=( src/spindle/configs/qwen35_9b_lora_16k.py src/spindle/configs/qwen35_9b_lora_64k.py + src/spindle/configs/qwen35_9b_lora_128k.py src/spindle/configs/qwen35_4b_fft_64k.py src/spindle/configs/gpt_oss_20b_lora_64k.py src/spindle/configs/qwen36_35b_a3b_lora_32k.py ) -spindle deploy "${deployment_files[@]}" "$@" +uv run --python 3.12 spindle deploy "${deployment_files[@]}" "$@" diff --git a/src/spindle/control_plane/deployments.py b/src/spindle/control_plane/deployments.py deleted file mode 100644 index cd3bb33..0000000 --- a/src/spindle/control_plane/deployments.py +++ /dev/null @@ -1,35 +0,0 @@ -"""Select recipes in deployment order, or explicitly by definition ID.""" - - -class DeploymentRoutes: - def __init__(self, definitions): - self.definitions = tuple(definitions) - - def select(self, model, mode=None): - candidates = [ - d for d in self.definitions if mode is None or d.parameterization == mode - ] - for definition in candidates: - if definition.definition_id == model: - return definition - for definition in candidates: - if definition.model == model: - return definition - return None - - def capabilities(self): - selected = {} - for definition in self.definitions: - selected.setdefault( - (definition.model, definition.parameterization), definition - ) - contexts = {} - for (model, _), definition in selected.items(): - contexts[model] = min( - contexts.get(model, definition.max_context_length), - definition.max_context_length, - ) - return [ - {"model_name": model, "max_context_length": context} - for model, context in contexts.items() - ] diff --git a/src/spindle/control_plane/http.py b/src/spindle/control_plane/http.py index 07d6223..b96f182 100644 --- a/src/spindle/control_plane/http.py +++ b/src/spindle/control_plane/http.py @@ -28,7 +28,6 @@ from spindle.request_timing import mark from spindle.telemetry.trainer import CommandMiddleware -from .deployments import DeploymentRoutes from .service import ControlPlane, FutureResolutionStatus ERROR_STATUSES: tuple[tuple[type[Exception], int, str], ...] = ( @@ -163,17 +162,15 @@ def create_control_plane_app( checkpoint_volume: str = "spindle-checkpoints", ) -> FastAPI: definitions = tuple(definitions) - - routes = DeploymentRoutes(definitions) - - def definition_for(model_name, parameterization): - selected = routes.select(model_name, parameterization) - return selected.definition_id if selected else None + definitions_by_name = { + definition.definition_id: definition for definition in definitions + } + for definition in definitions: + definitions_by_name.setdefault(definition.model, definition) + models = tuple(dict.fromkeys(definition.model for definition in definitions)) def supports_model(model_name: str) -> bool: - return any(definition.model == model_name for definition in definitions) or any( - definition.definition_id == model_name for definition in definitions - ) + return model_name in definitions_by_name async def authorize(request: Request) -> None: if api_key is not None and request.headers.get("x-api-key") != api_key: @@ -262,7 +259,15 @@ async def healthz() -> dict[str, str]: @app.get("/api/v1/get_server_capabilities") async def get_server_capabilities() -> dict[str, object]: - return {"supported_models": routes.capabilities()} + return { + "supported_models": [ + { + "model_name": model, + "max_context_length": definitions_by_name[model].max_context_length, + } + for model in models + ] + } @app.get("/api/v1/spindle/deployments") async def list_deployments(): @@ -336,8 +341,8 @@ async def create_model(body: CreateModelBody) -> dict[str, object]: raise ValueError("lora parameterization is configured through lora_config") if body.rollout is not None and parameterization != "full": raise ValueError("rollout is only configurable for full parameterization") - definition_id = definition_for(body.base_model, parameterization) - if definition_id is None: + definition = definitions_by_name.get(body.base_model) + if definition is None: raise HTTPException( status_code=400, detail=( @@ -345,14 +350,21 @@ async def create_model(body: CreateModelBody) -> dict[str, object]: f"{body.base_model}" ), ) + if definition.parameterization != parameterization: + raise HTTPException( + status_code=400, + detail=( + f"{definition.definition_id} uses " + f"{definition.parameterization} parameterization, not " + f"{parameterization}" + ), + ) creation = await control_plane.create_model( session_id=body.session_id, model_seq_id=body.model_seq_id, - definition_id=definition_id, + definition_id=definition.definition_id, spec={ - "base_model": next( - d.model for d in definitions if d.definition_id == definition_id - ), + "base_model": definition.model, "lora_config": body.lora_config, "parameterization": {"type": parameterization}, "rollout": body.rollout and body.rollout.model_dump(exclude_none=True), @@ -373,9 +385,8 @@ async def create_sampling_session( status_code=400, detail="base_model or model_path is required", ) - selected = routes.select(body.base_model) if body.base_model else None - definition_id = selected.definition_id if selected else None - if body.model_path is None and definition_id is None: + selected = definitions_by_name.get(body.base_model) if body.base_model else None + if body.model_path is None and selected is None: raise HTTPException( status_code=400, detail=f"unsupported base_model: {body.base_model}", @@ -383,13 +394,9 @@ async def create_sampling_session( session = await control_plane.create_sampling_session( session_id=body.session_id, sampling_session_seq_id=body.sampling_session_seq_id, - base_model=( - next(d.model for d in definitions if d.definition_id == definition_id) - if definition_id - else body.base_model - ), + base_model=selected.model if selected else body.base_model, model_path=body.model_path, - engine_definition_id=definition_id, + engine_definition_id=selected.definition_id if selected else None, ) if not supports_model(session.base_model): raise HTTPException( diff --git a/src/spindle/control_plane/records.py b/src/spindle/control_plane/records.py index 6528ec2..636eed4 100644 --- a/src/spindle/control_plane/records.py +++ b/src/spindle/control_plane/records.py @@ -62,22 +62,13 @@ class PlacementRecord(DurableRecord): engine_boot_id: str = "" -class SamplingSessionCreationRecord(DurableRecord): - session_id: Identifier - sampling_session_seq_id: int = Field(ge=0) - sampling_session_id: Identifier - fingerprint: NonEmptyString - created_at: Timestamp - session: dict | None = None - - class SamplingSessionRecord(DurableRecord): telemetry_tags: dict[str, str] = Field(default_factory=dict) sampling_session_id: Identifier session_id: Identifier sampling_session_seq_id: int = Field(ge=0) base_model: NonEmptyString - engine_definition_id: Identifier | None = None + engine_definition_id: Identifier model_path: str | None = None model_id: ModelIdentifier | None = None publish_version: PublishVersion | None = None @@ -87,6 +78,15 @@ class SamplingSessionRecord(DurableRecord): created_at: Timestamp +class SamplingSessionCreationRecord(DurableRecord): + session_id: Identifier + sampling_session_seq_id: int = Field(ge=0) + sampling_session_id: Identifier + fingerprint: NonEmptyString + created_at: Timestamp + session: SamplingSessionRecord + + class SamplerExportSubmissionRecord(DurableRecord): model_id: ModelIdentifier seq_id: int = Field(gt=0) @@ -110,7 +110,7 @@ class SamplerArtifactRecord(DurableRecord): model_id: ModelIdentifier export_seq_id: int = Field(gt=0) base_model: NonEmptyString - engine_definition_id: Identifier | None = None + engine_definition_id: Identifier publish_version: PublishVersion created_at: Timestamp expires_at: Timestamp | None = None diff --git a/src/spindle/control_plane/service.py b/src/spindle/control_plane/service.py index e1186df..91f009f 100644 --- a/src/spindle/control_plane/service.py +++ b/src/spindle/control_plane/service.py @@ -10,6 +10,7 @@ from datetime import UTC, datetime from enum import StrEnum from pathlib import PurePosixPath +from typing import cast from spindle.encoding import fingerprint from spindle.engine.api import EngineApi, FutureStatus @@ -61,6 +62,8 @@ SessionLastSeenRecord, SessionRecord, ) +from .store import typed_kv +from .trainer_reconciler import reconcile_trainers def request_id_for(model_id: str, seq_id: int) -> str: @@ -132,15 +135,13 @@ def __init__( sampling_task_stores: SessionKeyValueStores | None = None, read_checkpoint_metadata: Callable[[str], Awaitable[Mapping[str, object]]] | None = None, - reconcile_trainers: Callable[[str], Awaitable[bool | None]] | None = None, prepare_model: Callable[[ModelRecord], Awaitable[None]] | None = None, - trainer_autoscaling: Callable[[str], bool] = lambda _: False, list_checkpoints: CheckpointListing | None = None, delete_checkpoint: Callable[[str], Awaitable[None]] | None = None, checkpoint_root: str = "/checkpoints", creation_error: Callable[[str], Awaitable[str | None]] | None = None, ) -> None: - self.kv = kv + self.kv = typed_kv(kv) self.engines = engines self.sampling_tasks = sampling_tasks self.session_idle_timeout = session_idle_timeout @@ -149,9 +150,7 @@ def __init__( self.ensure_sampling_pool = ensure_sampling_pool self.sampling_task_stores = sampling_task_stores self.read_checkpoint_metadata = read_checkpoint_metadata - self.reconcile_trainers = reconcile_trainers self.prepare_model = prepare_model - self.trainer_autoscaling = trainer_autoscaling self.list_checkpoints = list_checkpoints self.delete_checkpoint = delete_checkpoint self.checkpoint_root = checkpoint_root @@ -186,7 +185,7 @@ async def create_session( return session async def heartbeat(self, session_id: str) -> SessionLastSeenRecord: - return await self._touch_session(session_id) + return await self._mark_session_active(session_id) async def close_session( self, @@ -205,7 +204,7 @@ async def close_session( closed.model_dump(mode="json"), ) await self._unload_session_models(session_id) - return SessionClosedRecord.model_validate(inserted.value) + return cast(SessionClosedRecord, inserted.value) async def create_model( self, @@ -216,7 +215,9 @@ async def create_model( spec: dict[str, object], model_id: str | None = None, ) -> ModelCreation: - await self._touch_session(session_id) + await self._mark_session_active(session_id) + + # hash request to ensure idempotency mark = fingerprint( "create_model", { @@ -224,7 +225,7 @@ async def create_model( "spec": spec, }, ) - anchor = ModelCreationRecord( + creation_record = ModelCreationRecord( session_id=session_id, model_seq_id=model_seq_id, model_id=model_id or self._model_id(session_id, model_seq_id), @@ -233,9 +234,9 @@ async def create_model( ) inserted = await self.kv.put_if_absent( model_creation_key(session_id, model_seq_id), - anchor.model_dump(mode="json"), + creation_record.model_dump(mode="json"), ) - stored = ModelCreationRecord.model_validate(inserted.value) + stored = cast(ModelCreationRecord, inserted.value) if stored.fingerprint != mark: raise SequenceConflict(session_id, model_seq_id) @@ -254,7 +255,7 @@ async def create_model( model_key(model.model_id), model.model_dump(mode="json"), ) - model = ModelRecord.model_validate(model_insert.value) + model = cast(ModelRecord, model_insert.value) if model_insert.created: await self.kv.put( trainer_demand_key(model.model_id), @@ -264,8 +265,7 @@ async def create_model( "created_at": model.created_at, }, ) - if self.reconcile_trainers is not None: - await self.reconcile_trainers(definition_id) + await self._trainer_demand_changed(definition_id) return ModelCreation( model, request_id_for(model.model_id, 0), @@ -466,14 +466,14 @@ async def submit_sampler_export(self, request: dict) -> str: ): raise ValueError("ttl_seconds must be a positive integer") model = await self.get_model(model_id) - await self._open_session(model.session_id) + await self._require_open_session(model.session_id) payload = { "path": name, "sampling_session_seq_id": sampling_session_seq_id, "ttl_seconds": ttl_seconds, } mark = fingerprint("save_weights_for_sampler", payload) - anchor = SamplerExportSubmissionRecord( + submission_record = SamplerExportSubmissionRecord( model_id=model_id, seq_id=seq_id, name=name, @@ -484,9 +484,9 @@ async def submit_sampler_export(self, request: dict) -> str: ) inserted = await self.kv.put_if_absent( sampler_export_submission_key(model_id, seq_id), - anchor.model_dump(mode="json"), + submission_record.model_dump(mode="json"), ) - stored = SamplerExportSubmissionRecord.model_validate(inserted.value) + stored = cast(SamplerExportSubmissionRecord, inserted.value) if stored.fingerprint != mark: sdk_retry = ( stored.name is None @@ -522,7 +522,7 @@ async def create_sampling_session( model_path: str | None = None, engine_definition_id: str | None = None, ) -> SamplingSessionRecord: - await self._touch_session(session_id) + await self._mark_session_active(session_id) mark = fingerprint( "create_sampling_session", {"base_model": base_model, "model_path": model_path}, @@ -530,30 +530,24 @@ async def create_sampling_session( key = sampling_session_creation_key(session_id, sampling_session_seq_id) value = await self.kv.get(key) if value is not None: - stored = SamplingSessionCreationRecord.model_validate(value) - if stored.session is not None and stored.fingerprint != mark: + stored = cast(SamplingSessionCreationRecord, value) + if stored.fingerprint != mark: raise SequenceConflict(session_id, sampling_session_seq_id) persisted = await self.kv.get( sampling_session_key(stored.sampling_session_id) ) - if persisted is not None: - session = SamplingSessionRecord.model_validate(persisted) - if stored.session is None and ( - model_path != session.model_path - or (base_model is not None and base_model != session.base_model) - ): - raise SequenceConflict(session_id, sampling_session_seq_id) - await self._ensure_sampling_pool(session) - return session - if stored.session is not None: - session = SamplingSessionRecord.model_validate(stored.session) - await self._validate_sampling_session(session) + session = stored.session + await self._validate_sampling_session(session) + if persisted is None: await self.kv.put( - sampling_session_key(session.sampling_session_id), - session.model_dump(mode="json"), + sampling_session_key(stored.sampling_session_id), + session, ) - await self._ensure_sampling_pool(session) - return session + else: + if cast(SamplingSessionRecord, persisted) != session: + raise SequenceConflict(session_id, sampling_session_seq_id) + await self._ensure_sampling_pool(session) + return session model_id = None publish_version = None export_seq_id = None @@ -567,12 +561,11 @@ async def create_sampling_session( raise ValueError("base_model does not match model_path") if ( engine_definition_id is not None - and artifact.engine_definition_id is not None and engine_definition_id != artifact.engine_definition_id ): raise ValueError("engine_definition_id does not match model_path") base_model = artifact.base_model - engine_definition_id = artifact.engine_definition_id or engine_definition_id + engine_definition_id = artifact.engine_definition_id model_id = artifact.model_id publish_version = artifact.publish_version export_seq_id = artifact.export_seq_id @@ -585,6 +578,8 @@ async def create_sampling_session( ) if not base_model: raise ValueError("base_model or model_path is required") + if not engine_definition_id: + raise ValueError("engine_definition_id is required") sampling_session_id = self._sampling_session_id( session_id, sampling_session_seq_id, @@ -604,38 +599,28 @@ async def create_sampling_session( expires_at=expires_at, created_at=self.clock(), ) - anchor = SamplingSessionCreationRecord( + creation_record = SamplingSessionCreationRecord( session_id=session_id, sampling_session_seq_id=sampling_session_seq_id, sampling_session_id=sampling_session_id, fingerprint=mark, created_at=session.created_at, - session=session.model_dump(mode="json"), + session=session, ) inserted = await self.kv.put_if_absent( key, - anchor.model_dump(mode="json"), - ) - stored = SamplingSessionCreationRecord.model_validate(inserted.value) - legacy_mark = fingerprint( - "create_sampling_session", - { - "base_model": base_model, - "model_path": model_path, - "model_id": model_id, - "publish_version": publish_version, - }, + creation_record.model_dump(mode="json"), ) - if stored.fingerprint not in {mark, legacy_mark}: + stored = cast(SamplingSessionCreationRecord, inserted.value) + if stored.fingerprint != mark: raise SequenceConflict(session_id, sampling_session_seq_id) - if stored.session is not None: - session = SamplingSessionRecord.model_validate(stored.session) + session = stored.session await self._validate_sampling_session(session) result = await self.kv.put_if_absent( sampling_session_key(session.sampling_session_id), - session.model_dump(mode="json"), + session, ) - persisted = SamplingSessionRecord.model_validate(result.value) + persisted = cast(SamplingSessionRecord, result.value) if persisted != session: raise SequenceConflict(session_id, sampling_session_seq_id) await self._ensure_sampling_pool(persisted) @@ -648,7 +633,7 @@ async def get_sampler_artifact( value = await self.kv.get(sampler_artifact_key(model_path)) if value is None: raise RecordNotFound("sampler artifact", model_path) - artifact = SamplerArtifactRecord.model_validate(value) + artifact = cast(SamplerArtifactRecord, value) if artifact.model_path != model_path: raise RecordNotFound("sampler artifact", model_path) self._check_expiry("sampler artifact", model_path, artifact.expires_at) @@ -662,7 +647,7 @@ async def _validate_sampling_session(self, session: SamplingSessionRecord) -> No self._check_expiry( "sampling session", session.sampling_session_id, session.expires_at ) - await self._open_session(session.session_id) + await self._require_open_session(session.session_id) async def get_sampling_session( self, @@ -671,7 +656,7 @@ async def get_sampling_session( value = await self.kv.get(sampling_session_key(sampling_session_id)) if value is None: raise RecordNotFound("sampling session", sampling_session_id) - session = SamplingSessionRecord.model_validate(value) + session = cast(SamplingSessionRecord, value) await self._validate_sampling_session(session) return session @@ -686,7 +671,7 @@ async def submit_sample(self, request: dict) -> str: sampling_session_id = str(request["sampling_session_id"]) seq_id = int(request["seq_id"]) session = await self.get_sampling_session(sampling_session_id) - await self._touch_session(session.session_id) + await self._mark_session_active(session.session_id) mark = fingerprint("sample", request) request_id = request_id_for(sampling_session_id, seq_id) key = sample_task_key(sampling_session_id, seq_id) @@ -702,7 +687,7 @@ async def submit_sample(self, request: dict) -> str: key, candidate.model_dump(mode="json"), ) - stored = SampleTaskRecord.model_validate(inserted.value) + stored = cast(SampleTaskRecord, inserted.value) if stored.fingerprint != mark: raise SequenceConflict(sampling_session_id, seq_id) if stored.task_id is not None: @@ -732,7 +717,7 @@ async def submit_sample(self, request: dict) -> str: async def engine_for(self, model_id: str) -> EngineApi: model = await self.get_model(model_id) - await self._touch_session(model.session_id) + await self._mark_session_active(model.session_id) placement = await self._placement(model_id) if placement is None: if await self._lost(model_id): @@ -743,8 +728,7 @@ async def engine_for(self, model_id: str) -> EngineApi: async def unload_model(self, model_id: str) -> str: model = await self._unload_model(model_id) - if self.reconcile_trainers is not None: - await self.reconcile_trainers(model.engine_definition_id) + await self._trainer_demand_changed(model.engine_definition_id) return f"{model.model_id}:unload" async def _unload_model(self, model_id: str) -> ModelRecord: @@ -871,7 +855,7 @@ async def _sampler_export_submission( value = await self.kv.get(sampler_export_submission_key(model_id, seq_id)) if value is None: return None - return SamplerExportSubmissionRecord.model_validate(value) + return cast(SamplerExportSubmissionRecord, value) async def _completed_sampler_export( self, @@ -881,7 +865,7 @@ async def _completed_sampler_export( path = self._sampler_model_path(export.model_id, export.name) value = await self.kv.get(sampler_artifact_key(path)) if value is not None: - artifact = SamplerArtifactRecord.model_validate(value) + artifact = cast(SamplerArtifactRecord, value) if ( artifact.model_path == path and artifact.model_id == export.model_id @@ -898,7 +882,7 @@ async def _completed_sampler_export( ) value = await self.kv.get(sampling_session_key(sampling_session_id)) if value is not None: - session = SamplingSessionRecord.model_validate(value) + session = cast(SamplingSessionRecord, value) if ( session.model_id == export.model_id and session.export_seq_id == export.seq_id @@ -913,7 +897,7 @@ async def _completed_sampler_export( ) if value is None: return None - receipt = SamplerExportResultRecord.model_validate(value) + receipt = cast(SamplerExportResultRecord, value) if receipt.model_id != export.model_id or receipt.seq_id != export.seq_id: return None model = await self.get_model(export.model_id) @@ -941,7 +925,7 @@ async def _persist_sampler_export_result( sampler_export_result_key(export.model_id, export.seq_id), receipt.model_dump(mode="json"), ) - stored = SamplerExportResultRecord.model_validate(inserted.value) + stored = cast(SamplerExportResultRecord, inserted.value) if ( stored.model_id != export.model_id or stored.seq_id != export.seq_id @@ -989,7 +973,7 @@ async def _finalize_sampler_export( ), expires_at, ) - await self._open_session(model.session_id) + await self._require_open_session(model.session_id) latest_path = self._sampler_model_path(model.model_id, "latest") latest_version_path = self._latest_sampler_model_path( model.model_id, @@ -1016,7 +1000,7 @@ async def _finalize_sampler_export( sampler_artifact_key(latest_version_path), versioned_latest.model_dump(mode="json"), ) - stored_latest = SamplerArtifactRecord.model_validate(inserted.value) + stored_latest = cast(SamplerArtifactRecord, inserted.value) if stored_latest.model_copy( update={"export_seq_id": export.seq_id, "created_at": completed_at} ).model_dump(exclude={"telemetry_tags"}) != versioned_latest.model_dump( @@ -1040,7 +1024,7 @@ async def _finalize_sampler_export( sampler_artifact_key(model_path), artifact.model_dump(mode="json"), ) - if SamplerArtifactRecord.model_validate(inserted.value).model_dump( + if cast(SamplerArtifactRecord, inserted.value).model_dump( exclude={"telemetry_tags"} ) != artifact.model_dump(exclude={"telemetry_tags"}): raise SequenceConflict(model.model_id, export.seq_id) @@ -1072,7 +1056,7 @@ async def _finalize_sampler_export( sampling_session_key(sampling_session_id), session.model_dump(mode="json"), ) - if SamplingSessionRecord.model_validate(inserted.value).model_dump( + if cast(SamplingSessionRecord, inserted.value).model_dump( exclude={"telemetry_tags"} ) != session.model_dump(exclude={"telemetry_tags"}): raise SequenceConflict( @@ -1098,7 +1082,7 @@ async def _retrieve_sample( ) if value is None: raise RecordNotFound("future", request_id) - record = SampleTaskRecord.model_validate(value) + record = cast(SampleTaskRecord, value) if record.request_id != request_id: raise RecordNotFound("future", request_id) if record.task_id is None: @@ -1133,7 +1117,7 @@ async def _retrieve_sample( def _sampling_task_store(self, session_id: str) -> KeyValueStore: if self.sampling_task_stores is None: return self.kv - return self.sampling_task_stores.for_session(session_id) + return typed_kv(self.sampling_task_stores.for_session(session_id)) async def _retrieve_creation( self, @@ -1192,7 +1176,7 @@ async def _lost(self, model_id: str) -> bool: async def _place(self, model: ModelRecord) -> PlacementRecord | None: value = await self.kv.get(placement_key(model.model_id)) if value is not None: - return PlacementRecord.model_validate(value) + return cast(PlacementRecord, value) if await self._lost(model.model_id): raise ModelLost(model.model_id) claim_key = placement_claim_key(model.model_id) @@ -1212,18 +1196,40 @@ async def _place(self, model: ModelRecord) -> PlacementRecord | None: if await self.kv.get(claim_key) == claim: await self.kv.delete(claim_key) + async def _trainer_demand_changed(self, _definition_id: str) -> None: + return None + + async def _reconcile_trainers( + self, + definition_id: str, + ) -> bool | None: + await reconcile_trainers( + self.kv, + self.engines, + definition_id, + revision=None, + minimum_instances=0, + maximum_instances=1, + models_per_instance=1, + ) + return None + + def _trainer_saturation_is_error(self, _definition_id: str) -> bool: + return False + async def _place_claimed(self, model: ModelRecord) -> PlacementRecord | None: definition_id = model.engine_definition_id try: active = await self.engines.active_instances(definition_id) instances = [instance for instance in active if instance.state == "running"] if not active: - if self.trainer_autoscaling(definition_id): - if self.reconcile_trainers is not None: - await self.reconcile_trainers(definition_id) + await self._reconcile_trainers(definition_id) + active = await self.engines.active_instances(definition_id) + instances = [ + instance for instance in active if instance.state == "running" + ] + if not active: return None - instance = await self.engines.ensure_instance(definition_id) - instances = [instance] if instance.state == "running" else [] instances.sort( key=lambda instance: hashlib.sha256( f"{model.model_id}\0{instance.instance_id}".encode() @@ -1257,9 +1263,9 @@ async def _place_claimed(self, model: ModelRecord) -> PlacementRecord | None: accepted_instance = instance break has_capacity = None - if accepted_instance is None and self.reconcile_trainers is not None: - has_capacity = await self.reconcile_trainers(definition_id) - if self.trainer_autoscaling(definition_id) and has_capacity: + if accepted_instance is None: + has_capacity = await self._reconcile_trainers(definition_id) + if has_capacity: return None if accepted_instance is None and self.session_idle_timeout is not None: for instance in instances: @@ -1287,7 +1293,7 @@ async def _place_claimed(self, model: ModelRecord) -> PlacementRecord | None: break if accepted_instance is None: if ( - self.trainer_autoscaling(definition_id) + self._trainer_saturation_is_error(definition_id) and has_capacity is False and all_instances_full ): @@ -1304,7 +1310,7 @@ async def _place_claimed(self, model: ModelRecord) -> PlacementRecord | None: placement_key(model.model_id), record.model_dump(mode="json"), ) - placement = PlacementRecord.model_validate(inserted.value) + placement = cast(PlacementRecord, inserted.value) if placement.engine_instance_id != accepted_instance.instance_id: try: await self.engines.client(accepted_instance.instance_id).unload_model( @@ -1334,13 +1340,13 @@ async def _reclaim_idle_models( for model_id in model_ids: value = await self.kv.get(model_key(model_id)) if value is not None: - incumbent = ModelRecord.model_validate(value) + incumbent = cast(ModelRecord, value) closed = await self.kv.get(session_closed_key(incumbent.session_id)) last_seen = await self.kv.get( session_last_seen_key(incumbent.session_id) ) seen_at = ( - SessionLastSeenRecord.model_validate(last_seen).seen_at + cast(SessionLastSeenRecord, last_seen).seen_at if last_seen is not None else 0.0 ) @@ -1385,26 +1391,24 @@ async def sweep_idle_sessions(self, idle_timeout: float) -> tuple[str, ...]: closed_sessions: set[str] = set() for key, value in session_items: if key.startswith("session_last_seen:"): - record = SessionLastSeenRecord.model_validate(value) + record = cast(SessionLastSeenRecord, value) last_seen[record.session_id] = record.seen_at elif key.startswith("session_closed:"): - closed_sessions.add( - SessionClosedRecord.model_validate(value).session_id - ) + closed_sessions.add(cast(SessionClosedRecord, value).session_id) elif key.startswith("session:"): - sessions.append(SessionRecord.model_validate(value).session_id) + sessions.append(cast(SessionRecord, value).session_id) models_by_session: dict[str, list[ModelRecord]] = defaultdict(list) creations_by_session: dict[str, list[str]] = defaultdict(list) placed: set[str] = set() for key, value in model_items: if key.startswith("model_creation:"): - creation = ModelCreationRecord.model_validate(value) + creation = cast(ModelCreationRecord, value) creations_by_session[creation.session_id].append(key) elif key.startswith("placement:"): - placed.add(PlacementRecord.model_validate(value).model_id) + placed.add(cast(PlacementRecord, value).model_id) elif key.startswith("model:"): - model = ModelRecord.model_validate(value) + model = cast(ModelRecord, value) models_by_session[model.session_id].append(model) closed: list[str] = [] @@ -1465,7 +1469,7 @@ async def sweep_idle_models(self, idle_timeout: float) -> tuple[str, ...]: async def sweep_idle_engines(self) -> tuple[str, ...]: items = await self.kv.list_items("trainer_demand:", "placement:") placed = { - PlacementRecord.model_validate(value).model_id + cast(PlacementRecord, value).model_id for key, value in items if key.startswith("placement:") } @@ -1502,13 +1506,12 @@ async def sweep_idle_engines(self) -> tuple[str, ...]: async def _unload_session_models(self, session_id: str) -> None: definitions = set() for _, value in await self.kv.list_items("model:"): - model = ModelRecord.model_validate(value) + model = cast(ModelRecord, value) if model.session_id == session_id: definitions.add(model.engine_definition_id) await self._unload_quietly(model.model_id) for definition_id in definitions: - if self.reconcile_trainers is not None: - await self.reconcile_trainers(definition_id) + await self._trainer_demand_changed(definition_id) async def _unload_quietly(self, model_id: str) -> None: try: @@ -1516,8 +1519,8 @@ async def _unload_quietly(self, model_id: str) -> None: except Exception: logging.getLogger(__name__).exception("unload %s", model_id) - async def _touch_session(self, session_id: str) -> SessionLastSeenRecord: - await self._open_session(session_id) + async def _mark_session_active(self, session_id: str) -> SessionLastSeenRecord: + await self._require_open_session(session_id) last_seen = SessionLastSeenRecord( session_id=session_id, seen_at=self.clock(), @@ -1528,25 +1531,25 @@ async def _touch_session(self, session_id: str) -> SessionLastSeenRecord: ) return last_seen - async def _open_session(self, session_id: str) -> SessionRecord: + async def _require_open_session(self, session_id: str) -> SessionRecord: if await self.kv.get(session_closed_key(session_id)) is not None: raise RecordUnavailable("session", session_id, "closed") value = await self.kv.get(session_key(session_id)) if value is None: raise RecordNotFound("session", session_id) - return SessionRecord.model_validate(value) + return cast(SessionRecord, value) async def get_model(self, model_id: str) -> ModelRecord: value = await self.kv.get(model_key(model_id)) if value is None: raise RecordNotFound("model", model_id) - return ModelRecord.model_validate(value) + return cast(ModelRecord, value) async def _placement(self, model_id: str) -> PlacementRecord | None: value = await self.kv.get(placement_key(model_id)) if value is None: return None - return PlacementRecord.model_validate(value) + return cast(PlacementRecord, value) async def _live_instance(self, placement: PlacementRecord) -> EngineInstance: instance = await self.engines.get_instance(placement.engine_instance_id) diff --git a/src/spindle/control_plane/store.py b/src/spindle/control_plane/store.py new file mode 100644 index 0000000..66e91c2 --- /dev/null +++ b/src/spindle/control_plane/store.py @@ -0,0 +1,87 @@ +from __future__ import annotations + +from pydantic import BaseModel + +from spindle.providers.contracts import InsertResult, KeyValueStore + +from .records import ( + ModelCreationRecord, + ModelRecord, + PlacementRecord, + SamplerArtifactRecord, + SamplerExportResultRecord, + SamplerExportSubmissionRecord, + SampleTaskRecord, + SamplingSessionCreationRecord, + SamplingSessionRecord, + SessionClosedRecord, + SessionLastSeenRecord, + SessionRecord, +) + +RECORD_TYPES: dict[str, type[BaseModel]] = { + "session": SessionRecord, + "session_last_seen": SessionLastSeenRecord, + "session_closed": SessionClosedRecord, + "model_creation": ModelCreationRecord, + "model": ModelRecord, + "placement": PlacementRecord, + "sampling_session_creation": SamplingSessionCreationRecord, + "sampling_session": SamplingSessionRecord, + "sampler_export_submission": SamplerExportSubmissionRecord, + "sampler_export_result": SamplerExportResultRecord, + "sampler_artifact": SamplerArtifactRecord, + "sample_task": SampleTaskRecord, +} + + +class TypedKeyValueStore: + """Validate durable control-plane records at the KV boundary.""" + + def __init__(self, store: KeyValueStore) -> None: + self.store = store + + def _record_type(self, key: str) -> type[BaseModel] | None: + return RECORD_TYPES.get(key.partition(":")[0]) + + def _decode(self, key: str, value: object) -> object: + record_type = self._record_type(key) + return record_type.model_validate(value) if record_type is not None else value + + def _encode(self, key: str, value: object) -> object: + record_type = self._record_type(key) + if record_type is None: + return ( + value.model_dump(mode="json") if isinstance(value, BaseModel) else value + ) + return record_type.model_validate(value).model_dump(mode="json") + + async def get(self, key: str) -> object | None: + value = await self.store.get(key) + return None if value is None else self._decode(key, value) + + async def put(self, key: str, value: object) -> None: + await self.store.put(key, self._encode(key, value)) + + async def put_if_absent(self, key: str, value: object) -> InsertResult: + result = await self.store.put_if_absent(key, self._encode(key, value)) + return InsertResult(result.created, self._decode(key, result.value)) + + async def delete(self, key: str) -> None: + await self.store.delete(key) + + async def list_keys(self, prefix: str) -> tuple[str, ...]: + return await self.store.list_keys(prefix) + + async def list_items( + self, + *prefixes: str, + ) -> tuple[tuple[str, object], ...]: + return tuple( + (key, self._decode(key, value)) + for key, value in await self.store.list_items(*prefixes) + ) + + +def typed_kv(store: KeyValueStore) -> TypedKeyValueStore: + return store if isinstance(store, TypedKeyValueStore) else TypedKeyValueStore(store) diff --git a/src/spindle/control_plane/trainer_reconciler.py b/src/spindle/control_plane/trainer_reconciler.py new file mode 100644 index 0000000..f8790b2 --- /dev/null +++ b/src/spindle/control_plane/trainer_reconciler.py @@ -0,0 +1,248 @@ +from __future__ import annotations + +import logging +import math +import time +from collections.abc import Callable +from typing import cast + +from pydantic import BaseModel + +from spindle.control_plane.records import ( + ModelRecord, + PlacementRecord, + SessionClosedRecord, +) +from spindle.providers.contracts import ( + EngineInstance, + EnginePlatform, + KeyValueStore, +) + +from .store import typed_kv + +TRAINER_PLAN_PREFIX = "trainer_plan:" + +_log = logging.getLogger(__name__) + + +class TrainerPlanRecord(BaseModel): + definition_id: str + observed_at: float + live_models: int + desired_instances: int + maximum_instances: int | None + active_instances: int + starting_instances: int + draining_instances: int + pending_models: int + converged: bool + updated_at: float + + +def trainer_plan_key(definition_id: str) -> str: + return f"{TRAINER_PLAN_PREFIX}{definition_id}" + + +async def reconcile_trainers( + kv: KeyValueStore, + engines: EnginePlatform, + definition_id: str, + *, + revision: str | None, + maximum_instances: int | None, + minimum_instances: int = 0, + models_per_instance: int = 1, + scale_up: bool = True, + clock: Callable[[], float] = time.time, +) -> TrainerPlanRecord: + if minimum_instances < 0: + raise ValueError("minimum_instances must be non-negative") + if maximum_instances is not None and maximum_instances < minimum_instances: + raise ValueError("maximum_instances must be at least minimum_instances") + if models_per_instance < 1: + raise ValueError("models_per_instance must be positive") + kv = typed_kv(kv) + + def current_revision(instance: EngineInstance) -> bool: + return revision is None or instance.revision == revision + + instances = await engines.list_instances() + items = await kv.list_items( + "model:", + "placement:", + "session_closed:", + "trainer_demand:", + ) + closed = { + cast(SessionClosedRecord, value).session_id + for key, value in items + if key.startswith("session_closed:") + } + models = [ + cast(ModelRecord, value) + for key, value in items + if key.startswith("model:") + and cast(ModelRecord, value).engine_definition_id == definition_id + and cast(ModelRecord, value).session_id not in closed + ] + placements = { + record.model_id: record + for key, value in items + if key.startswith("placement:") + for record in (cast(PlacementRecord, value),) + } + demand_ids = { + str(value["model_id"]) + for key, value in items + if key.startswith("trainer_demand:") + and isinstance(value, dict) + and value.get("definition_id") == definition_id + and isinstance(value.get("model_id"), str) + } | set(placements) + records = [ + instance + for instance in instances + if instance.definition_id == definition_id and not instance.terminal + ] + current_ids = {record.instance_id for record in records if current_revision(record)} + demanded_models = [model for model in models if model.model_id in demand_ids] + live_models = [ + model + for model in demanded_models + if ( + model.model_id not in placements + or placements[model.model_id].engine_instance_id in current_ids + ) + ] + desired = max( + minimum_instances, + math.ceil(len(live_models) / models_per_instance), + ) + if maximum_instances is not None: + desired = min(desired, maximum_instances) + current = [record for record in records if current_revision(record)] + usable = [record for record in current if record.state in {"starting", "running"}] + if scale_up: + missing = max(0, desired - len(usable)) + if ( + missing + and maximum_instances is not None + and len(records) >= maximum_instances + ): + stale = [ + record + for record in records + if not current_revision(record) + and record.state in {"running", "draining"} + ] + await _stop_empty( + engines, + stale, + len(records) - maximum_instances + missing, + ) + records = [ + instance + for instance in await engines.list_instances() + if instance.definition_id == definition_id and not instance.terminal + ] + available = ( + missing + if maximum_instances is None + else max(0, maximum_instances - len(records)) + ) + for _ in range(min(max(0, desired - len(usable)), available)): + instance = await engines.spawn_instance(definition_id) + usable.append(instance) + records = [ + instance + for instance in await engines.list_instances() + if instance.definition_id == definition_id and not instance.terminal + ] + current = [record for record in records if current_revision(record)] + running = sorted( + (record for record in current if record.state == "running"), + key=lambda record: record.instance_id, + ) + await _stop_empty(engines, running, max(0, len(current) - desired)) + for record in records: + if not current_revision(record) and record.state == "running": + await engines.set_instance_state(record.instance_id, "draining") + draining = [ + instance + for instance in await engines.list_instances() + if instance.definition_id == definition_id and instance.state == "draining" + ] + await _stop_empty(engines, draining, len(draining)) + final = [ + instance + for instance in await engines.list_instances() + if instance.definition_id == definition_id and not instance.terminal + ] + current_final = [record for record in final if current_revision(record)] + active = sum(record.state == "running" for record in current_final) + starting = sum(record.state == "starting" for record in current_final) + draining_count = sum(record.state == "draining" for record in final) + stale_starting = any( + record.state == "starting" and not current_revision(record) for record in final + ) + plan = TrainerPlanRecord( + definition_id=definition_id, + observed_at=clock(), + live_models=len(demanded_models), + desired_instances=desired, + maximum_instances=maximum_instances, + active_instances=active, + starting_instances=starting, + draining_instances=draining_count, + pending_models=( + 0 + if maximum_instances is None + else max( + 0, + len(live_models) - maximum_instances * models_per_instance, + ) + ), + converged=( + active + starting == desired and draining_count == 0 and not stale_starting + ), + updated_at=clock(), + ) + await kv.put(trainer_plan_key(definition_id), plan.model_dump(mode="json")) + return plan + + +async def _model_ids( + engines: EnginePlatform, + records: list[EngineInstance], +) -> dict[str, tuple[str, ...]]: + result = {} + for record in records: + try: + result[record.instance_id] = await engines.client( + record.instance_id + ).model_ids() + except Exception: # noqa: BLE001 - one unreachable engine must not block others + _log.exception("list models on %s", record.instance_id) + return result + + +async def _stop_empty( + engines: EnginePlatform, + records: list[EngineInstance], + limit: int, +) -> None: + infos = await _model_ids(engines, records) + stopped = 0 + for record in records: + if stopped >= limit or infos.get(record.instance_id) != (): + continue + try: + acknowledged = await engines.client(record.instance_id).shutdown_if_idle() + except Exception: # noqa: BLE001 - failed shutdowns are retried later + _log.exception("shutdown %s", record.instance_id) + continue + if not acknowledged: + continue + await engines.stop_instance(record.instance_id) + stopped += 1 diff --git a/src/spindle/providers/contracts/sampling.py b/src/spindle/providers/contracts/sampling.py index 2631009..ecd7eaa 100644 --- a/src/spindle/providers/contracts/sampling.py +++ b/src/spindle/providers/contracts/sampling.py @@ -11,7 +11,7 @@ class SamplingTask: session_id: str sampling_session_id: str base_model: str - engine_definition_id: str | None + engine_definition_id: str model_path: str | None model_id: str | None publish_version: int | None diff --git a/src/spindle/providers/local/engines.py b/src/spindle/providers/local/engines.py index 2c236c9..20e8c9a 100644 --- a/src/spindle/providers/local/engines.py +++ b/src/spindle/providers/local/engines.py @@ -37,15 +37,25 @@ async def ensure_instance( if definition_id != self.definition_id: raise RecordNotFound("engine definition", definition_id) async with self._lock: - current = ( - self._instances.get(self._current) if self._current else None - ) + current = self._instances.get(self._current) if self._current else None if ( current is not None and current.state in {"starting", "running"} and current.revision == self.revision ): return current + current = next( + ( + instance + for instance in self._instances.values() + if instance.state in {"starting", "running"} + and instance.revision == self.revision + ), + None, + ) + if current is not None: + self._current = current.instance_id + return current instance = self._spawn_instance() self._current = instance.instance_id return instance diff --git a/src/spindle/providers/modal/app.py b/src/spindle/providers/modal/app.py index b4097c7..e31c81b 100644 --- a/src/spindle/providers/modal/app.py +++ b/src/spindle/providers/modal/app.py @@ -10,15 +10,13 @@ from huggingface_hub import snapshot_download from stitch.pools.modal_flash import ModalFlashPool -from spindle.control_plane import ControlPlane, create_control_plane_app +from spindle.control_plane import create_control_plane_app from spindle.control_plane.keys import model_key, placement_key, trainer_demand_key from spindle.control_plane.records import ModelRecord +from spindle.control_plane.trainer_reconciler import reconcile_trainers from spindle.deployments import validate_frontend from spindle.inference.sampling import sample_task -from spindle.providers.contracts import ( - Parameterization, - SamplingTask, -) +from spindle.providers.contracts import SamplingTask from spindle.telemetry.otlp import sample_trace from .checkpoint_storage import ( @@ -32,6 +30,7 @@ configs_from_env, platform_from_env, ) +from .deployment_control import DeploymentControlPlane from .engines import ModalEnginePlatform from .fft_pool import ( FFTPoolSpec, @@ -68,13 +67,13 @@ from .trainer_reconciler import ( complete_reconcile, pending_reconciliations, - reconcile_trainers, release_reconcile_call, request_reconcile, ) DEFINITIONS = tuple(configs_from_env()) validate_frontend([definition.recipe for definition in DEFINITIONS]) +DEFINITIONS_BY_ID = {definition.definition_id: definition for definition in DEFINITIONS} SETTINGS = DEFINITIONS[0] PLATFORM = platform_from_env() APP_NAME = PLATFORM["frontend"] @@ -102,24 +101,6 @@ app = modal.App(APP_NAME) -async def _read_checkpoint_metadata(uri: str) -> dict[str, object]: - return await ModalCheckpointStorage( - checkpoint_volume, CHECKPOINT_ROOT, lock=CHECKPOINT_READ_LOCK - ).read_metadata(uri) - - -async def _list_checkpoints(model_id: str | None) -> list[dict[str, object]]: - return await ModalCheckpointStorage( - checkpoint_volume, CHECKPOINT_ROOT, lock=CHECKPOINT_READ_LOCK - ).list(model_id) - - -async def _delete_checkpoint(uri: str) -> None: - await ModalCheckpointStorage( - checkpoint_volume, CHECKPOINT_ROOT, lock=CHECKPOINT_READ_LOCK - ).delete(uri) - - image = ( modal.Image.debian_slim(python_version="3.12") .apt_install("git") @@ -276,10 +257,10 @@ async def execute_sample(task: dict) -> dict: async def _execute_sample(task: dict, stats: dict) -> dict: definition_id = str(task["engine_definition_id"]) - parameterization = parameterization_for(definition_id) + definition = module_for(definition_id) + parameterization = definition.parameterization if parameterization not in {"full", "lora"}: raise ValueError(f"unsupported sampling definition: {definition_id}") - definition = module_for(definition_id) rollout_world_size = definition.recipe.inference_gpus_per_node rollout_tensor_parallel_size = definition.rollout_tensor_parallel_size if rollout_world_size % rollout_tensor_parallel_size: @@ -340,21 +321,18 @@ async def _model_record(kv, model_id: str): def module_for(definition_id: str): - for definition in DEFINITIONS: - if definition.definition_id == definition_id: - return definition - raise KeyError(definition_id) - - -def parameterization_for(definition_id: str) -> Parameterization | None: try: - return module_for(definition_id).parameterization + return DEFINITIONS_BY_ID[definition_id] except KeyError: - return None + raise KeyError(f"definition is not deployed: {definition_id}") from None + +def trainer_maximum_instances(definition_id: str) -> int: + return module_for(definition_id).recipe.trainer_max_instances -def trainer_autoscaling(definition_id: str) -> bool: - return parameterization_for(definition_id) is not None + +def trainer_models_per_instance(definition_id: str) -> int: + return module_for(definition_id).recipe.trainer_max_clients_per_instance @app.function( @@ -370,21 +348,20 @@ async def trainer_reconciler(delay_seconds: float = 0.0) -> None: call_id = modal.current_function_call_id() async def run(definition_id: str, token: str) -> None: - parameterization = parameterization_for(definition_id) - if parameterization is None or await deployment_error(definition_id): + definition = DEFINITIONS_BY_ID.get(definition_id) + if definition is None or await deployment_error(definition_id): await complete_reconcile(definition_id, token) return - module = module_for(definition_id) - maximum_instances = module.recipe.trainer_max_instances try: await reconcile_trainers( shared_kv(), ModalEnginePlatform(shared_kv(), _spawn_engine), definition_id, revision=None, - maximum_instances=maximum_instances, - models_per_instance=module.recipe.trainer_max_clients_per_instance, - scale_up=trainer_autoscaling(definition_id), + minimum_instances=0, + maximum_instances=trainer_maximum_instances(definition_id), + models_per_instance=trainer_models_per_instance(definition_id), + scale_up=True, ) except Exception: logging.getLogger(__name__).exception( @@ -407,8 +384,7 @@ async def run(definition_id: str, token: str) -> None: async def kick_trainer_reconciler(definition_id: str) -> None: - if parameterization_for(definition_id) is None: - return + module_for(definition_id) async def spawn(delay_seconds: float) -> str: call = await trainer_reconciler.spawn.aio(delay_seconds) @@ -451,20 +427,45 @@ async def clear_deployment_failure(definition_id: str) -> None: await kick_trainer_reconciler(definition_id) -def _plane(): - kv = shared_kv() - task_stores = ModalSessionKeyValueStores() - engines = ModalEnginePlatform(kv, _spawn_engine) +class ModalDeploymentControlPlane(DeploymentControlPlane): + def __init__(self): + kv = shared_kv() + task_stores = ModalSessionKeyValueStores() + engines = ModalEnginePlatform(kv, _spawn_engine) + storage = ModalCheckpointStorage( + checkpoint_volume, + CHECKPOINT_ROOT, + lock=CHECKPOINT_READ_LOCK, + ) + super().__init__( + kv, + engines, + sampling_tasks=ModalSamplingTaskPlatform( + task_stores, + self._spawn_sampling, + ), + session_idle_timeout=SESSION_IDLE_TIMEOUT, + ensure_sampling_pool=self._ensure_deployment_pool, + prepare_model=self._prepare_deployment_model, + creation_error=deployment_error, + sampling_task_stores=task_stores, + read_checkpoint_metadata=storage.read_metadata, + list_checkpoints=storage.list, + delete_checkpoint=storage.delete, + checkpoint_root=CHECKPOINT_ROOT, + request_trainer_reconciliation=self._request_trainer_reconciliation, + trainer_maximum_instances=trainer_maximum_instances, + trainer_models_per_instance=trainer_models_per_instance, + ) - async def spawn_sampling(task: SamplingTask) -> str: + async def _spawn_sampling(self, task: SamplingTask) -> str: call = await execute_sample.spawn.aio(asdict(task)) return call.object_id - async def prepare_model(model) -> None: - parameterization = parameterization_for(model.engine_definition_id) - if parameterization is None: - return - await prepare_model_assets.remote.aio(model.engine_definition_id) + async def _prepare_deployment_model(self, model: ModelRecord) -> None: + definition = module_for(model.engine_definition_id) + await prepare_model_assets.remote.aio(definition.definition_id) + parameterization = definition.parameterization if parameterization == "full": await ensure_fft_pool.spawn.aio(_latest_pool(model).as_dict()) else: @@ -472,9 +473,9 @@ async def prepare_model(model) -> None: LoraPoolSpec(model.engine_definition_id).as_dict() ) - async def ensure_pool(session) -> None: + async def _ensure_deployment_pool(self, session) -> None: definition_id = session.engine_definition_id - parameterization = parameterization_for(definition_id) + parameterization = module_for(definition_id).parameterization if parameterization == "lora": if session.model_id is None: await prepare_model_assets.remote.aio(definition_id) @@ -500,41 +501,34 @@ async def ensure_pool(session) -> None: if session.model_id is None: await prepare_model_assets.remote.aio(definition_id) if pool.latest: - pool = _latest_pool(await _model_record(kv, session.model_id)) + pool = _latest_pool(await _model_record(self.kv, session.model_id)) await ensure_fft_pool.remote.aio(pool.as_dict()) else: await _touch_fft_pool(pool) - async def kick_trainers(definition_id: str) -> bool: + async def _request_trainer_reconciliation( + self, + definition_id: str, + *, + maximum_instances: int | None, + **_: object, + ) -> bool: + module_for(definition_id) await kick_trainer_reconciler(definition_id) - if not trainer_autoscaling(definition_id): - return False - maximum = module_for(definition_id).recipe.trainer_max_instances instances = [ instance - for instance in await engines.list_instances() + for instance in await self.engines.list_instances() if instance.definition_id == definition_id and not instance.terminal ] - return any( - instance.state in {"starting", "draining"} for instance in instances - ) or len(instances) < int(maximum) - - return ControlPlane( - kv, - engines, - sampling_tasks=ModalSamplingTaskPlatform(task_stores, spawn_sampling), - session_idle_timeout=SESSION_IDLE_TIMEOUT, - ensure_sampling_pool=ensure_pool, - prepare_model=prepare_model, - creation_error=deployment_error, - sampling_task_stores=task_stores, - read_checkpoint_metadata=_read_checkpoint_metadata, - list_checkpoints=_list_checkpoints, - delete_checkpoint=_delete_checkpoint, - checkpoint_root=CHECKPOINT_ROOT, - reconcile_trainers=kick_trainers, - trainer_autoscaling=trainer_autoscaling, - ) + return ( + any(instance.state in {"starting", "draining"} for instance in instances) + or maximum_instances is None + or len(instances) < maximum_instances + ) + + +def build_deployment_control_plane() -> ModalDeploymentControlPlane: + return ModalDeploymentControlPlane() @app.function( @@ -550,7 +544,7 @@ async def kick_trainers(definition_id: str) -> bool: @modal.asgi_app(requires_proxy_auth=False) def server(): return create_control_plane_app( - _plane(), + build_deployment_control_plane(), DEFINITIONS, api_key=os.environ["TINKER_API_KEY"], checkpoint_volume=CHECKPOINT_VOLUME_NAME, @@ -562,7 +556,7 @@ async def _lose_undefined_models() -> tuple[str, ...]: lost = [] for _, value in await kv.list_items("model:"): model = ModelRecord.model_validate(value) - if parameterization_for(model.engine_definition_id) is not None: + if model.engine_definition_id in DEFINITIONS_BY_ID: continue await kv.delete(placement_key(model.model_id)) await kv.delete(trainer_demand_key(model.model_id)) @@ -580,7 +574,8 @@ async def _cleanup_fft_pools() -> tuple[str, ...]: ).app_name for _, value in await shared_kv().list_items("model:") for model in (ModelRecord.model_validate(value),) - if parameterization_for(model.engine_definition_id) == "full" + for definition in (DEFINITIONS_BY_ID.get(model.engine_definition_id),) + if definition is not None and definition.parameterization == "full" } stopped = [] registry = fft_pool_kv() @@ -624,7 +619,8 @@ async def _cleanup_lora_pools() -> tuple[str, ...]: LoraPoolSpec(model.engine_definition_id).app_name for _, value in await registry.list_items("model:") for model in (ModelRecord.model_validate(value),) - if parameterization_for(model.engine_definition_id) == "lora" + for definition in (DEFINITIONS_BY_ID.get(model.engine_definition_id),) + if definition is not None and definition.parameterization == "lora" } stopped = [] for key, value in await registry.list_items("lora_pool:"): @@ -648,7 +644,7 @@ async def _cleanup_lora_pools() -> tuple[str, ...]: @app.function(image=image, env=TRAINER_DEPLOYMENT_ENV, schedule=SWEEP_PERIOD) def cleaner(): async def run() -> None: - plane = _plane() + plane = build_deployment_control_plane() await plane.sweep_idle_sessions(SESSION_IDLE_TIMEOUT) await plane.sweep_idle_models(SESSION_IDLE_TIMEOUT) await plane.sweep_idle_engines() diff --git a/src/spindle/providers/modal/deployment_control.py b/src/spindle/providers/modal/deployment_control.py new file mode 100644 index 0000000..8ed017a --- /dev/null +++ b/src/spindle/providers/modal/deployment_control.py @@ -0,0 +1,35 @@ +from collections.abc import Awaitable, Callable + +from spindle.control_plane import ControlPlane + + +class DeploymentControlPlane(ControlPlane): + def __init__( + self, + *args, + request_trainer_reconciliation: Callable[..., Awaitable[object]], + trainer_maximum_instances: Callable[[str], int], + trainer_models_per_instance: Callable[[str], int], + **kwargs, + ): + super().__init__(*args, **kwargs) + self.request_trainer_reconciliation = request_trainer_reconciliation + self.trainer_maximum_instances = trainer_maximum_instances + self.trainer_models_per_instance = trainer_models_per_instance + + async def _trainer_demand_changed(self, definition_id: str) -> None: + await self._reconcile_trainers(definition_id) + + async def _reconcile_trainers(self, definition_id: str) -> bool | None: + result = await self.request_trainer_reconciliation( + definition_id, + revision=None, + minimum_instances=0, + maximum_instances=self.trainer_maximum_instances(definition_id), + models_per_instance=self.trainer_models_per_instance(definition_id), + scale_up=True, + ) + return result if isinstance(result, bool) else None + + def _trainer_saturation_is_error(self, _definition_id: str) -> bool: + return True diff --git a/src/spindle/providers/modal/kv.py b/src/spindle/providers/modal/kv.py index b569f40..1ddd771 100644 --- a/src/spindle/providers/modal/kv.py +++ b/src/spindle/providers/modal/kv.py @@ -7,6 +7,9 @@ import modal from grpclib.exceptions import StreamTerminatedError +from pydantic import BaseModel + +from spindle.control_plane.store import RECORD_TYPES from ..contracts import InsertResult @@ -69,20 +72,43 @@ def app_store_name(name: str, app_id: str) -> str: class ModalKeyValueStore: - def __init__(self, dictionary: modal.Dict) -> None: + def __init__( + self, + dictionary: modal.Dict, + record_types: dict[str, type[BaseModel]] = RECORD_TYPES, + ) -> None: self.dictionary = dictionary + self.record_types = record_types + + def _record_type(self, key: str) -> type[BaseModel] | None: + return self.record_types.get(key.partition(":")[0]) + + def _decode(self, key: str, value: object) -> object: + record_type = self._record_type(key) + return record_type.model_validate(value) if record_type is not None else value + + def _encode(self, key: str, value: object) -> object: + record_type = self._record_type(key) + if record_type is None: + return ( + value.model_dump(mode="json") if isinstance(value, BaseModel) else value + ) + return record_type.model_validate(value).model_dump(mode="json") async def get(self, key: str) -> object | None: - return await self.dictionary.get.aio(key) + value = await self.dictionary.get.aio(key) + return None if value is None else self._decode(key, value) async def put(self, key: str, value: object) -> None: - await self.dictionary.put.aio(key, value) + await self.dictionary.put.aio(key, self._encode(key, value)) async def put_if_absent(self, key: str, value: object) -> InsertResult: - created = await self.dictionary.put.aio(key, value, skip_if_exists=True) + stored = self._encode(key, value) + created = await self.dictionary.put.aio(key, stored, skip_if_exists=True) if created: - return InsertResult(True, value) - return InsertResult(False, await self.dictionary.get.aio(key)) + return InsertResult(True, self._decode(key, stored)) + existing = await self.dictionary.get.aio(key) + return InsertResult(False, self._decode(key, existing)) async def delete(self, key: str) -> None: await self.dictionary.pop.aio(key, None) @@ -98,7 +124,7 @@ async def list_items( try: async with asyncio.timeout(LIST_ITEMS_TIMEOUT_SECONDS): items = [ - (key, value) + (key, self._decode(key, value)) async for key, value in self.dictionary.items.aio() if any(key.startswith(prefix) for prefix in prefixes) ] diff --git a/src/spindle/providers/modal/scoped.py b/src/spindle/providers/modal/scoped.py index fc113f0..9f7b4ef 100644 --- a/src/spindle/providers/modal/scoped.py +++ b/src/spindle/providers/modal/scoped.py @@ -422,6 +422,8 @@ async def gateway(): @modal.concurrent(max_inputs=128) @modal.asgi_app(requires_proxy_auth=False) def api(): + engines = ModalEnginePlatform(shared_kv(), spawn_engine) + async def prepare_model(model): if model.spec.get("rollout"): raise ValueError( @@ -440,7 +442,7 @@ async def spawn_sampling(task): storage = ModalCheckpointStorage(checkpoints) plane = ScopedControlPlane( shared_kv(), - ModalEnginePlatform(shared_kv(), spawn_engine), + engines, session_idle_timeout=None, prepare_model=prepare_model, ensure_sampling_pool=ensure_pool, diff --git a/src/spindle/providers/modal/scoped_assignment.py b/src/spindle/providers/modal/scoped_assignment.py index 3ef99be..c990850 100644 --- a/src/spindle/providers/modal/scoped_assignment.py +++ b/src/spindle/providers/modal/scoped_assignment.py @@ -41,6 +41,4 @@ async def claim_model(registry, kv, engines, definition_id, model_id): await registry.put.aio("model:" + model_id, route) # This single assignment is the sampling admission fence and store selector. await registry.put.aio("slot:0", model_id) - if not active: - await engines.spawn_instance(definition_id) return await registry.get.aio("model:" + model_id) diff --git a/src/spindle/providers/modal/scoped_control.py b/src/spindle/providers/modal/scoped_control.py index c794c20..104f406 100644 --- a/src/spindle/providers/modal/scoped_control.py +++ b/src/spindle/providers/modal/scoped_control.py @@ -1,5 +1,7 @@ """Retry-safe model preparation for scoped deployments.""" + from spindle.control_plane import ControlPlane +from spindle.control_plane.trainer_reconciler import reconcile_trainers class ScopedControlPlane(ControlPlane): @@ -13,3 +15,17 @@ async def create_model(self, **kwargs): # request must finish preparation, not silently skip a failed first try. await self.prepare_scoped_model(creation.model) return creation + + async def _trainer_demand_changed(self, definition_id: str) -> None: + await self._reconcile_trainers(definition_id) + + async def _reconcile_trainers(self, definition_id: str) -> None: + await reconcile_trainers( + self.kv, + self.engines, + definition_id, + revision=None, + minimum_instances=0, + maximum_instances=1, + models_per_instance=1, + ) diff --git a/src/spindle/providers/modal/trainer_reconciler.py b/src/spindle/providers/modal/trainer_reconciler.py index 27f8ac8..51d9721 100644 --- a/src/spindle/providers/modal/trainer_reconciler.py +++ b/src/spindle/providers/modal/trainer_reconciler.py @@ -1,24 +1,13 @@ from __future__ import annotations import logging -import math import time import uuid from collections.abc import Awaitable, Callable import modal -from pydantic import BaseModel -from spindle.control_plane.records import ( - ModelRecord, - PlacementRecord, - SessionClosedRecord, -) -from spindle.providers.contracts import ( - EngineInstance, - EnginePlatform, - KeyValueStore, -) +from spindle.control_plane.trainer_reconciler import TrainerPlanRecord from .kv import shared_kv @@ -33,20 +22,6 @@ _log = logging.getLogger(__name__) -class TrainerPlanRecord(BaseModel): - definition_id: str - observed_at: float - live_models: int - desired_instances: int - maximum_instances: int | None - active_instances: int - starting_instances: int - draining_instances: int - pending_models: int - converged: bool - updated_at: float - - def trainer_plan_key(definition_id: str) -> str: return f"{TRAINER_PLAN_PREFIX}{definition_id}" @@ -136,199 +111,6 @@ async def list_trainer_plans() -> tuple[TrainerPlanRecord, ...]: ) -async def reconcile_trainers( - kv: KeyValueStore, - engines: EnginePlatform, - definition_id: str, - *, - revision: str | None, - maximum_instances: int | None, - models_per_instance: int = 1, - scale_up: bool = True, - clock: Callable[[], float] = time.time, -) -> TrainerPlanRecord: - def current_revision(instance: EngineInstance) -> bool: - return revision is None or instance.revision == revision - - instances = await engines.list_instances() - items = await kv.list_items( - "model:", - "placement:", - "session_closed:", - "trainer_demand:", - ) - closed = { - SessionClosedRecord.model_validate(value).session_id - for key, value in items - if key.startswith("session_closed:") - } - models = [ - ModelRecord.model_validate(value) - for key, value in items - if key.startswith("model:") - and value["engine_definition_id"] == definition_id - and value["session_id"] not in closed - ] - placements = { - record.model_id: record - for key, value in items - if key.startswith("placement:") - for record in (PlacementRecord.model_validate(value),) - } - demand_ids = { - str(value["model_id"]) - for key, value in items - if key.startswith("trainer_demand:") - and isinstance(value, dict) - and value.get("definition_id") == definition_id - and isinstance(value.get("model_id"), str) - } | set(placements) - records = [ - instance - for instance in instances - if instance.definition_id == definition_id and not instance.terminal - ] - current_ids = {record.instance_id for record in records if current_revision(record)} - demanded_models = [model for model in models if model.model_id in demand_ids] - live_models = [ - model - for model in demanded_models - if ( - model.model_id not in placements - or placements[model.model_id].engine_instance_id in current_ids - ) - ] - desired = math.ceil(len(live_models) / models_per_instance) - if maximum_instances is not None: - desired = min(desired, maximum_instances) - current = [record for record in records if current_revision(record)] - usable = [record for record in current if record.state in {"starting", "running"}] - if scale_up: - missing = max(0, desired - len(usable)) - if ( - missing - and maximum_instances is not None - and len(records) >= maximum_instances - ): - stale = [ - record - for record in records - if not current_revision(record) - and record.state in {"running", "draining"} - ] - await _stop_empty( - engines, - stale, - len(records) - maximum_instances + missing, - ) - records = [ - instance - for instance in await engines.list_instances() - if instance.definition_id == definition_id - and not instance.terminal - ] - available = ( - missing - if maximum_instances is None - else max(0, maximum_instances - len(records)) - ) - for _ in range(min(max(0, desired - len(usable)), available)): - instance = await engines.spawn_instance(definition_id) - usable.append(instance) - records = [ - instance - for instance in await engines.list_instances() - if instance.definition_id == definition_id and not instance.terminal - ] - current = [record for record in records if current_revision(record)] - running = sorted( - (record for record in current if record.state == "running"), - key=lambda record: record.instance_id, - ) - await _stop_empty(engines, running, max(0, len(current) - desired)) - for record in records: - if not current_revision(record) and record.state == "running": - await engines.set_instance_state(record.instance_id, "draining") - draining = [ - instance - for instance in await engines.list_instances() - if instance.definition_id == definition_id and instance.state == "draining" - ] - await _stop_empty(engines, draining, len(draining)) - final = [ - instance - for instance in await engines.list_instances() - if instance.definition_id == definition_id and not instance.terminal - ] - current_final = [record for record in final if current_revision(record)] - active = sum(record.state == "running" for record in current_final) - starting = sum(record.state == "starting" for record in current_final) - draining_count = sum(record.state == "draining" for record in final) - stale_starting = any( - record.state == "starting" and not current_revision(record) for record in final - ) - plan = TrainerPlanRecord( - definition_id=definition_id, - observed_at=clock(), - live_models=len(demanded_models), - desired_instances=desired, - maximum_instances=maximum_instances, - active_instances=active, - starting_instances=starting, - draining_instances=draining_count, - pending_models=( - 0 - if maximum_instances is None - else max( - 0, - len(live_models) - maximum_instances * models_per_instance, - ) - ), - converged=( - active + starting == desired and draining_count == 0 and not stale_starting - ), - updated_at=clock(), - ) - await kv.put(trainer_plan_key(definition_id), plan.model_dump(mode="json")) - return plan - - -async def _model_ids( - engines: EnginePlatform, - records: list[EngineInstance], -) -> dict[str, tuple[str, ...]]: - result = {} - for record in records: - try: - result[record.instance_id] = await engines.client( - record.instance_id - ).model_ids() - except Exception: - _log.exception("list models on %s", record.instance_id) - return result - - -async def _stop_empty( - engines: EnginePlatform, - records: list[EngineInstance], - limit: int, -) -> None: - infos = await _model_ids(engines, records) - stopped = 0 - for record in records: - if stopped >= limit or infos.get(record.instance_id) != (): - continue - try: - acknowledged = await engines.client(record.instance_id).shutdown_if_idle() - except Exception: - _log.exception("shutdown %s", record.instance_id) - continue - if not acknowledged: - continue - await engines.stop_instance(record.instance_id) - stopped += 1 - - async def _call_finished(call_id: str) -> bool: try: await modal.FunctionCall.from_id(call_id).get.aio(timeout=0) @@ -344,6 +126,6 @@ async def _call_finished(call_id: str) -> bool: ): _log.exception("trainer reconcile liveness poll %s", call_id) return False - except Exception: + except Exception: # noqa: BLE001 - any completed remote failure ends the call return True return True diff --git a/tests/control_plane/test_http.py b/tests/control_plane/test_http.py index 4b49992..9470808 100644 --- a/tests/control_plane/test_http.py +++ b/tests/control_plane/test_http.py @@ -359,7 +359,7 @@ async def run() -> None: json={ "session_id": session.json()["session_id"], "model_seq_id": 0, - "base_model": BASE_MODEL, + "base_model": f"{DEFINITION}_full", "parameterization": {"type": "full"}, }, ) @@ -433,31 +433,42 @@ async def run() -> None: client = http_client() session_id, _ = await created_model(client) - def create(seq: int, **body): + def create(seq: int, base_model: str = BASE_MODEL, **body): return client.post( "/api/v1/create_model", json={ "session_id": session_id, "model_seq_id": seq, - "base_model": BASE_MODEL, + "base_model": base_model, **body, }, ) full = {"parameterization": {"type": "full"}} accepted = await create( - 1, **full, rollout={"min_containers": 2, "max_containers": 4} + 1, + base_model=f"{DEFINITION}_full", + **full, + rollout={"min_containers": 2, "max_containers": 4}, ) assert accepted.status_code == 200 lora = await create(2, lora_config={"rank": 8}, rollout={"max_containers": 4}) assert lora.status_code == 400 assert "only configurable for full" in lora.json()["message"] inverted = await create( - 3, **full, rollout={"min_containers": 4, "max_containers": 2} + 3, + base_model=f"{DEFINITION}_full", + **full, + rollout={"min_containers": 4, "max_containers": 2}, ) assert inverted.status_code == 400 assert inverted.json()["error"] == "invalid_request" - unknown = await create(4, **full, rollout={"target_concurrency": 8}) + unknown = await create( + 4, + base_model=f"{DEFINITION}_full", + **full, + rollout={"target_concurrency": 8}, + ) assert unknown.status_code == 400 assert unknown.json()["error"] == "invalid_request" await client.aclose() diff --git a/tests/control_plane/test_models.py b/tests/control_plane/test_models.py index 44f50eb..829f6d0 100644 --- a/tests/control_plane/test_models.py +++ b/tests/control_plane/test_models.py @@ -2,7 +2,6 @@ import pytest -from tests.support import EchoExecutor from spindle.control_plane import ControlPlane, FutureResolutionStatus from spindle.control_plane.keys import model_key, placement_key, trainer_demand_key from spindle.engine import OperationKind @@ -15,6 +14,8 @@ InMemoryKeyValueStore, LocalEnginePlatform, ) +from spindle.providers.modal.deployment_control import DeploymentControlPlane +from tests.support import EchoExecutor DEFINITION = "qwen3_4b_lora32_16k" @@ -441,16 +442,19 @@ async def prepare(model) -> None: if attempts == 1: raise RuntimeError("asset download failed") - async def reconcile(definition_id: str) -> None: + async def reconcile(definition_id: str, **_: object) -> None: reconciled.append(definition_id) + class DemandControlPlane(ControlPlane): + async def _trainer_demand_changed(self, definition_id: str) -> None: + await reconcile(definition_id) + async def run() -> None: kv = InMemoryKeyValueStore() - plane = ControlPlane( + plane = DemandControlPlane( kv, LocalEnginePlatform(DEFINITION, EchoExecutor), prepare_model=prepare, - reconcile_trainers=reconcile, ) session = await plane.create_session() request = { @@ -475,18 +479,19 @@ async def run() -> None: def test_model_creation_reports_saturation_at_trainer_cap() -> None: - async def no_capacity(definition_id: str) -> bool: + async def no_capacity(_definition_id: str, **_: object) -> bool: return False async def run() -> None: kv = InMemoryKeyValueStore() engines = LocalEnginePlatform(DEFINITION, EchoExecutor, max_models=1) await engines.spawn_instance(DEFINITION) - plane = ControlPlane( + plane = DeploymentControlPlane( kv, engines, - reconcile_trainers=no_capacity, - trainer_autoscaling=lambda _: True, + request_trainer_reconciliation=no_capacity, + trainer_maximum_instances=lambda _: 1, + trainer_models_per_instance=lambda _: 1, session_idle_timeout=300, ) session = await plane.create_session() @@ -515,18 +520,19 @@ async def run() -> None: def test_model_creation_reuses_released_slot_at_trainer_cap() -> None: - async def no_capacity(definition_id: str) -> bool: + async def no_capacity(_definition_id: str, **_: object) -> bool: return False async def run() -> None: kv = InMemoryKeyValueStore() engines = LocalEnginePlatform(DEFINITION, EchoExecutor, max_models=2) instance = await engines.spawn_instance(DEFINITION) - plane = ControlPlane( + plane = DeploymentControlPlane( kv, engines, - reconcile_trainers=no_capacity, - trainer_autoscaling=lambda _: True, + request_trainer_reconciliation=no_capacity, + trainer_maximum_instances=lambda _: 1, + trainer_models_per_instance=lambda _: 2, ) session = await plane.create_session() creations = [ diff --git a/tests/control_plane/test_operations.py b/tests/control_plane/test_operations.py index a662566..3730bd8 100644 --- a/tests/control_plane/test_operations.py +++ b/tests/control_plane/test_operations.py @@ -66,9 +66,7 @@ async def run() -> None: "data": [ { "loss_fn_inputs": {}, - "model_input": { - "chunks": [{"tokens": [1]}] - }, + "model_input": {"chunks": [{"tokens": [1]}]}, } ], "loss_fn": "cross_entropy", @@ -85,9 +83,7 @@ async def run() -> None: later = await engine.optim_step( {"model_id": model_id, "seq_id": 2, "adam_params": {}} ) - assert ( - await plane.retrieve(later) - ).status == FutureResolutionStatus.PENDING + assert (await plane.retrieve(later)).status == FutureResolutionStatus.PENDING await engine.forward_backward( forward_backward_body(model_id, 1), "application/json", @@ -231,10 +227,10 @@ def test_placement_failure_leaves_creation_pending() -> None: class FlakyPlatform(LocalEnginePlatform): fail = True - async def ensure_instance(self, definition): + async def spawn_instance(self, definition): if self.fail: raise RuntimeError("platform outage") - return await super().ensure_instance(definition) + return await super().spawn_instance(definition) async def run() -> None: engines = FlakyPlatform(DEFINITION, EchoExecutor) diff --git a/tests/control_plane/test_sampler_exports.py b/tests/control_plane/test_sampler_exports.py index 490dcdc..f7992d0 100644 --- a/tests/control_plane/test_sampler_exports.py +++ b/tests/control_plane/test_sampler_exports.py @@ -2,7 +2,6 @@ import pytest -from tests.support import TinkerStubExecutor from spindle.control_plane import ControlPlane, FutureResolutionStatus from spindle.control_plane.keys import ( sampler_artifact_key, @@ -16,6 +15,7 @@ InMemoryKeyValueStore, LocalEnginePlatform, ) +from tests.support import TinkerStubExecutor BASE_MODEL = "Qwen/Qwen3-8B" DEFINITION = "qwen3_8b" @@ -100,9 +100,6 @@ async def run() -> None: assert latest.model_path == f"{latest_path}/000007" assert latest.latest assert latest.publish_version == 7 - legacy = sampling.model_dump(mode="json") - legacy.pop("engine_definition_id") - assert type(sampling).model_validate(legacy).engine_definition_id is None assert await plane.submit_sampler_export(request) == request_id assert (await plane.retrieve(request_id)).result == resolution.result @@ -232,8 +229,13 @@ async def run() -> None: sampling_session_seq_id=0, model_path=result.result["path"], ) - assert await plane.kv.get(sampler_artifact_key(result.result["path"])) is not None - assert await plane.kv.get(sampling_session_key(session.sampling_session_id)) is not None + assert ( + await plane.kv.get(sampler_artifact_key(result.result["path"])) is not None + ) + assert ( + await plane.kv.get(sampling_session_key(session.sampling_session_id)) + is not None + ) asyncio.run(run()) @@ -243,13 +245,15 @@ def test_expired_creation_retry_never_ensures_pool(missing_record) -> None: async def run() -> None: now = 100.0 plane, session_id, model_id = await plane_with_model(lambda: now) - request_id = await plane.submit_sampler_export(export_request(model_id, ttl_seconds=10)) - result = await plane.retrieve(request_id, timeout=1.0) - kwargs = dict( - session_id=session_id, - sampling_session_seq_id=0, - model_path=result.result["path"], + request_id = await plane.submit_sampler_export( + export_request(model_id, ttl_seconds=10) ) + result = await plane.retrieve(request_id, timeout=1.0) + kwargs = { + "session_id": session_id, + "sampling_session_seq_id": 0, + "model_path": result.result["path"], + } session = await plane.create_sampling_session(**kwargs) if missing_record: await plane.kv.delete(sampling_session_key(session.sampling_session_id)) @@ -262,7 +266,10 @@ async def unexpected_pool(_): with pytest.raises(RecordUnavailable, match="expired"): await plane.create_sampling_session(**kwargs) if missing_record: - assert await plane.kv.get(sampling_session_key(session.sampling_session_id)) is None + assert ( + await plane.kv.get(sampling_session_key(session.sampling_session_id)) + is None + ) asyncio.run(run()) @@ -283,7 +290,8 @@ async def run() -> None: result = await plane.retrieve(request_id, timeout=1.0) key = ( sampler_artifact_key(result.result["path"]) - if named else sampling_session_key(result.result["sampling_session_id"]) + if named + else sampling_session_key(result.result["sampling_session_id"]) ) if missing_record: await plane.kv.delete(key) @@ -330,7 +338,9 @@ def test_expired_named_artifact_remains_reserved() -> None: async def run() -> None: now = 100.0 plane, _, model_id = await plane_with_model(lambda: now) - request_id = await plane.submit_sampler_export(export_request(model_id, ttl_seconds=10)) + request_id = await plane.submit_sampler_export( + export_request(model_id, ttl_seconds=10) + ) result = await plane.retrieve(request_id, timeout=1.0) key = sampler_artifact_key(result.result["path"]) original = await plane.kv.get(key) diff --git a/tests/control_plane/test_sdk_e2e.py b/tests/control_plane/test_sdk_e2e.py index c31ebc1..1a6ea24 100644 --- a/tests/control_plane/test_sdk_e2e.py +++ b/tests/control_plane/test_sdk_e2e.py @@ -1,5 +1,4 @@ import asyncio -import importlib import itertools import json import os @@ -20,6 +19,7 @@ LocalEnginePlatform, LocalSamplingTaskPlatform, ) +from spindle.providers.modal.checkpoint_storage import ModalCheckpointStorage from tests.support import TinkerStubExecutor, TinkerStubSampler, serve BASE_MODEL = "Qwen/Qwen3-8B" @@ -338,11 +338,9 @@ def commit(self) -> None: return None -def volume_plane(tmp_path, monkeypatch) -> tuple[ControlPlane, object, list[str]]: - modal_app = importlib.import_module("spindle.providers.modal.app") +def volume_plane(tmp_path) -> tuple[ControlPlane, object, list[str]]: root = tmp_path / "checkpoints" - monkeypatch.setattr(modal_app, "CHECKPOINT_ROOT", str(root)) - monkeypatch.setattr(modal_app, "checkpoint_volume", FakeVolume()) + storage = ModalCheckpointStorage(FakeVolume(), str(root)) clock = itertools.count(1_700_000_000) loaded: list[str] = [] @@ -364,16 +362,16 @@ async def persist_checkpoint(self, model_id, payload, snapshot): plane = ControlPlane( InMemoryKeyValueStore(), LocalEnginePlatform(DEFINITION, VolumeExecutor), - read_checkpoint_metadata=modal_app._read_checkpoint_metadata, - list_checkpoints=modal_app._list_checkpoints, - delete_checkpoint=modal_app._delete_checkpoint, + read_checkpoint_metadata=storage.read_metadata, + list_checkpoints=storage.list, + delete_checkpoint=storage.delete, checkpoint_root=str(root), ) return plane, root, loaded -def test_real_sdk_lists_and_deletes_checkpoints(tmp_path, monkeypatch) -> None: - plane, root, loaded = volume_plane(tmp_path, monkeypatch) +def test_real_sdk_lists_and_deletes_checkpoints(tmp_path) -> None: + plane, root, loaded = volume_plane(tmp_path) app = create_control_plane_app( plane, DEFINITIONS, api_key=API_KEY, retrieve_window=5.0 ) @@ -445,8 +443,8 @@ def test_real_sdk_lists_and_deletes_checkpoints(tmp_path, monkeypatch) -> None: assert first.is_dir() -def test_real_sdk_lost_model_fails_fast(tmp_path, monkeypatch) -> None: - plane, _, _ = volume_plane(tmp_path, monkeypatch) +def test_real_sdk_lost_model_fails_fast(tmp_path) -> None: + plane, _, _ = volume_plane(tmp_path) app = create_control_plane_app( plane, DEFINITIONS, api_key=API_KEY, retrieve_window=5.0 ) diff --git a/tests/providers/test_checkpoint_storage.py b/tests/providers/test_checkpoint_storage.py index 59720ae..80699f5 100644 --- a/tests/providers/test_checkpoint_storage.py +++ b/tests/providers/test_checkpoint_storage.py @@ -10,6 +10,7 @@ from spindle.providers.modal.app import DEFINITIONS, PLATFORM from spindle.providers.modal.checkpoint_storage import ( CHECKPOINT_ROOT, + ModalCheckpointStorage, _scan_checkpoints, ) from spindle.providers.modal.deployment_apps import volumes_for @@ -77,6 +78,7 @@ def scan(model_id): patch.object(app.checkpoint_volume, "reload"), patch.object(app.checkpoint_volume, "commit"), ): - asyncio.run(app._delete_checkpoint(str(tmp_path / "final/run-a"))) + storage = ModalCheckpointStorage(app.checkpoint_volume, str(tmp_path)) + asyncio.run(storage.delete(str(tmp_path / "final/run-a"))) assert scan("run-a") == [] assert len(scan("run-b")) == 1 diff --git a/tests/providers/test_definition_registry.py b/tests/providers/test_definition_registry.py index c9bffe2..a6af38b 100644 --- a/tests/providers/test_definition_registry.py +++ b/tests/providers/test_definition_registry.py @@ -1,14 +1,10 @@ from spindle.providers.modal.app import ( DEFINITIONS, module_for, - parameterization_for, ) def test_definition_registry_resolves_every_definition() -> None: for definition in DEFINITIONS: assert module_for(definition.definition_id) is definition - assert parameterization_for(definition.definition_id) == ( - definition.parameterization - ) assert definition.trainer_app_name diff --git a/tests/providers/test_modal_app.py b/tests/providers/test_modal_app.py index f346d32..a7b1b32 100644 --- a/tests/providers/test_modal_app.py +++ b/tests/providers/test_modal_app.py @@ -40,11 +40,10 @@ def test_definitions_come_only_from_the_configured_records() -> None: } -def test_trainer_autoscaling_supports_full_and_lora_definitions() -> None: +def test_definition_lookup_rejects_missing_definition() -> None: modal_app = importlib.import_module("spindle.providers.modal.app") - assert modal_app.trainer_autoscaling(FULL_DEFINITION) - assert modal_app.trainer_autoscaling(LORA_DEFINITION) - assert not modal_app.trainer_autoscaling("missing-definition") + with pytest.raises(KeyError, match="definition is not deployed"): + modal_app.module_for("missing-definition") @pytest.mark.parametrize( @@ -74,8 +73,8 @@ async def kick(_definition_id: str) -> None: monkeypatch.setattr(modal_app, "kick_trainer_reconciler", kick) modal_app.module_for(LORA_DEFINITION).recipe.trainer_max_instances = 1 - plane = modal_app._plane() - assert asyncio.run(plane.reconcile_trainers(LORA_DEFINITION)) is available + plane = modal_app.build_deployment_control_plane() + assert asyncio.run(plane._reconcile_trainers(LORA_DEFINITION)) is available def test_ensure_pool_deploys_pinned_base_pool(monkeypatch) -> None: @@ -113,7 +112,7 @@ async def ensure(spec: dict) -> str: ) async def run() -> None: - plane = modal_app._plane() + plane = modal_app.build_deployment_control_plane() await plane.ensure_sampling_pool(session) await plane.ensure_sampling_pool(session) @@ -159,7 +158,7 @@ def model(definition_id: str, spec: dict) -> SimpleNamespace: ) async def run() -> None: - plane = modal_app._plane() + plane = modal_app.build_deployment_control_plane() await plane.prepare_model( model( FULL_DEFINITION, {"rollout": {"min_containers": 8, "max_containers": 8}} @@ -244,7 +243,7 @@ def session(latest: bool, version: int) -> SimpleNamespace: async def run() -> None: await kv.put(model_key(model_id), model.model_dump(mode="json")) - plane = modal_app._plane() + plane = modal_app.build_deployment_control_plane() await plane.ensure_sampling_pool(session(True, 0)) await plane.ensure_sampling_pool(session(True, 0)) await plane.ensure_sampling_pool(session(False, 3)) @@ -403,10 +402,11 @@ def reload(self) -> None: monkeypatch.setattr(modal_app, "checkpoint_volume", Volume()) checkpoint_uri = root / "model" / "weights" / "checkpoint" - assert ( - asyncio.run(modal_app._read_checkpoint_metadata(str(checkpoint_uri))) - == metadata + storage = modal_app.ModalCheckpointStorage( + modal_app.checkpoint_volume, + str(root), ) + assert asyncio.run(storage.read_metadata(str(checkpoint_uri))) == metadata assert reloads == [True] @@ -554,7 +554,8 @@ async def run() -> None: asyncio.run(run()) assert stopped == [pools["orphan"].app_name] - assert modal_app.parameterization_for("removed_definition") is None + with pytest.raises(KeyError, match="definition is not deployed"): + modal_app.module_for("removed_definition") def test_checkpoint_volume_listing_and_delete(tmp_path, monkeypatch) -> None: @@ -574,6 +575,10 @@ def commit(self) -> None: checkpoints = tmp_path / "checkpoints" monkeypatch.setattr(modal_app, "CHECKPOINT_ROOT", str(checkpoints)) monkeypatch.setattr(modal_app, "checkpoint_volume", Volume("ckpt")) + storage = modal_app.ModalCheckpointStorage( + modal_app.checkpoint_volume, + str(checkpoints), + ) lora = checkpoints / "step-1" / "model-a" lora.mkdir(parents=True) @@ -588,7 +593,7 @@ def commit(self) -> None: (checkpoints / "step-1" / "stray.txt").write_text("x") (checkpoints / "step-1" / "incomplete").mkdir() - entries = asyncio.run(modal_app._list_checkpoints(None)) + entries = asyncio.run(storage.list(None)) assert sorted( (entry["model_id"], entry["name"], entry["size_bytes"], entry["metadata"]) for entry in entries @@ -602,20 +607,20 @@ def commit(self) -> None: ("model-b", "latest", 12, {}), ] assert {entry["path"] for entry in entries} == {str(lora), str(fft)} - assert [ - entry["name"] for entry in asyncio.run(modal_app._list_checkpoints("model-b")) - ] == ["latest"] - assert asyncio.run(modal_app._list_checkpoints("model-c")) == [] + assert [entry["name"] for entry in asyncio.run(storage.list("model-b"))] == [ + "latest" + ] + assert asyncio.run(storage.list("model-c")) == [] assert events == ["reload:ckpt"] * 3 events.clear() - asyncio.run(modal_app._delete_checkpoint(str(lora))) + asyncio.run(storage.delete(str(lora))) assert not lora.exists() assert events == ["reload:ckpt", "commit:ckpt"] with pytest.raises(RecordNotFound): - asyncio.run(modal_app._delete_checkpoint(str(lora))) + asyncio.run(storage.delete(str(lora))) with pytest.raises(ValueError): - asyncio.run(modal_app._delete_checkpoint(str(tmp_path / "elsewhere"))) + asyncio.run(storage.delete(str(tmp_path / "elsewhere"))) def test_pool_cleanup_continues_after_failure_and_retries_entry(monkeypatch, caplog): @@ -833,7 +838,7 @@ async def deploy(record): ) async def run(): - plane = modal_app._plane() + plane = modal_app.build_deployment_control_plane() await asyncio.gather(*(plane.ensure_sampling_pool(session) for _ in range(6))) await modal_app._ready_lora_pool(LoraPoolSpec(LORA_DEFINITION)) diff --git a/tests/providers/test_modal_kv.py b/tests/providers/test_modal_kv.py index c0a295e..ad7747c 100644 --- a/tests/providers/test_modal_kv.py +++ b/tests/providers/test_modal_kv.py @@ -4,6 +4,7 @@ import pytest +from spindle.control_plane.records import SessionRecord from spindle.providers.local.kv import InMemoryKeyValueStore from spindle.providers.modal import kv @@ -28,6 +29,34 @@ class RetryableStreamError(Exception): pass +def test_modal_store_validates_typed_records() -> None: + async def run() -> None: + values = {} + + async def get(key): + return values.get(key) + + async def put(key, value, skip_if_exists=False): + if skip_if_exists and key in values: + return False + values[key] = value + return True + + store = kv.ModalKeyValueStore( + SimpleNamespace( + get=SimpleNamespace(aio=get), + put=SimpleNamespace(aio=put), + ) + ) + session = SessionRecord(session_id="session", created_at=1.0) + await store.put("session:session", session) + + assert values["session:session"] == session.model_dump(mode="json") + assert await store.get("session:session") == session + + asyncio.run(run()) + + class StreamingItems: def __init__(self, failures: int) -> None: self.failures = failures @@ -75,7 +104,7 @@ async def stream(): def test_list_items_retries_terminated_stream(monkeypatch) -> None: async def run() -> None: items = StreamingItems(failures=2) - store = kv.ModalKeyValueStore(SimpleNamespace(items=items)) + store = kv.ModalKeyValueStore(SimpleNamespace(items=items), record_types={}) sleep = AsyncMock() monkeypatch.setattr(kv, "StreamTerminatedError", RetryableStreamError) monkeypatch.setattr(kv.asyncio, "sleep", sleep) @@ -92,7 +121,7 @@ async def run() -> None: def test_list_items_discards_partial_retry(monkeypatch) -> None: async def run() -> None: items = PartiallyTerminatedItems() - store = kv.ModalKeyValueStore(SimpleNamespace(items=items)) + store = kv.ModalKeyValueStore(SimpleNamespace(items=items), record_types={}) monkeypatch.setattr(kv, "StreamTerminatedError", RetryableStreamError) monkeypatch.setattr(kv.asyncio, "sleep", AsyncMock()) @@ -108,7 +137,7 @@ async def run() -> None: def test_list_items_retries_stalled_stream(monkeypatch) -> None: async def run() -> None: items = HangingItems() - store = kv.ModalKeyValueStore(SimpleNamespace(items=items)) + store = kv.ModalKeyValueStore(SimpleNamespace(items=items), record_types={}) monkeypatch.setattr(kv, "LIST_ITEMS_TIMEOUT_SECONDS", 0.01) monkeypatch.setattr(kv.asyncio, "sleep", AsyncMock()) @@ -122,7 +151,7 @@ async def run() -> None: def test_list_items_does_not_retry_other_errors(monkeypatch) -> None: async def run() -> None: items = StreamingItems(failures=1) - store = kv.ModalKeyValueStore(SimpleNamespace(items=items)) + store = kv.ModalKeyValueStore(SimpleNamespace(items=items), record_types={}) monkeypatch.setattr(kv, "StreamTerminatedError", ValueError) sleep = AsyncMock() monkeypatch.setattr(kv.asyncio, "sleep", sleep) @@ -138,7 +167,7 @@ async def run() -> None: def test_list_items_reraises_after_retry_limit(monkeypatch) -> None: async def run() -> None: items = StreamingItems(failures=3) - store = kv.ModalKeyValueStore(SimpleNamespace(items=items)) + store = kv.ModalKeyValueStore(SimpleNamespace(items=items), record_types={}) monkeypatch.setattr(kv, "StreamTerminatedError", RetryableStreamError) monkeypatch.setattr(kv.asyncio, "sleep", AsyncMock()) diff --git a/tests/providers/test_trainer_reconciler.py b/tests/providers/test_trainer_reconciler.py index 3e619c7..fd48646 100644 --- a/tests/providers/test_trainer_reconciler.py +++ b/tests/providers/test_trainer_reconciler.py @@ -1,11 +1,11 @@ import asyncio from spindle.control_plane import ControlPlane, FutureResolutionStatus +from spindle.control_plane.trainer_reconciler import reconcile_trainers from spindle.providers.local import ( InMemoryKeyValueStore, LocalEnginePlatform, ) -from spindle.providers.modal.trainer_reconciler import reconcile_trainers from tests.support import EchoExecutor DEFINITION = "qwen_full" diff --git a/tests/scoped/test_replacement.py b/tests/scoped/test_replacement.py index c92c78d..6dd6b35 100644 --- a/tests/scoped/test_replacement.py +++ b/tests/scoped/test_replacement.py @@ -6,7 +6,12 @@ from stitch.sync import ConstraintUnmet from stitch.types import VersionConstraint, VersionRef -from spindle.inference.scoped_sidecar import AssignedReconciler, AssignedSnapshotStore, assigned_app, _expected_run +from spindle.inference.scoped_sidecar import ( + AssignedReconciler, + AssignedSnapshotStore, + _expected_run, + assigned_app, +) from spindle.providers.modal.scoped_assignment import claim_model @@ -15,138 +20,196 @@ def __init__(self, values): self.values = values self.get = SimpleNamespace(aio=self.read) self.put = SimpleNamespace(aio=self.write) - async def read(self, key): return self.values.get(key) - async def write(self, key, value): self.values[key] = value + + async def read(self, key): + return self.values.get(key) + + async def write(self, key, value): + self.values[key] = value def test_replacement_requires_confirmed_loss_and_fences_previous_model(): async def check(): - values = {'routes': [{}, {'url': 'latest'}], 'trainer_demand:a': {}, 'trainer_demand:b': {}} + values = { + "routes": [{}, {"url": "latest"}], + "trainer_demand:a": {}, + "trainer_demand:b": {}, + } registry = Registry(values) active = [] deleted = [] - async def instances(_): return list(active) - async def spawn(_): active.append('trainer') - async def delete(key): deleted.append(key) - engines = SimpleNamespace(active_instances=instances, spawn_instance=spawn) - async def get(key): return values.get(key) + + async def instances(_): + return list(active) + + async def delete(key): + deleted.append(key) + + engines = SimpleNamespace(active_instances=instances) + + async def get(key): + return values.get(key) + kv = SimpleNamespace(delete=delete, get=get) - await claim_model(registry, kv, engines, 'engine', 'a') - await claim_model(registry, kv, engines, 'engine', 'a') - assert len(active) == 1 - with pytest.raises(ValueError, match='already active'): - await claim_model(registry, kv, engines, 'engine', 'b') - assert values['slot:0'] == 'a' + await claim_model(registry, kv, engines, "engine", "a") + await claim_model(registry, kv, engines, "engine", "a") + assert active == [] + active.append("trainer") # The trainer reconciler fulfilled demand. + with pytest.raises(ValueError, match="already active"): + await claim_model(registry, kv, engines, "engine", "b") + assert values["slot:0"] == "a" active.clear() # Provider has confirmed the trainer invocation ended. - await claim_model(registry, kv, engines, 'engine', 'b') - assert values['slot:0'] == 'b' and deleted == ['trainer_demand:a'] - assert len(active) == 1 - with pytest.raises(ValueError, match='replaced'): - await claim_model(registry, kv, engines, 'engine', 'a') + await claim_model(registry, kv, engines, "engine", "b") + assert values["slot:0"] == "b" and deleted == ["trainer_demand:a"] + assert active == [] + with pytest.raises(ValueError, match="replaced"): + await claim_model(registry, kv, engines, "engine", "a") + asyncio.run(check()) -def test_retry_after_spawn_failure_keeps_assignment(): +def test_claim_retry_keeps_assignment_without_provisioning(): async def check(): - registry = Registry({'routes': [{}, {'url': 'latest'}], 'trainer_demand:a': {}}) - attempts = [] - async def instances(_): return [] - async def spawn(_): - attempts.append(1) - if len(attempts) == 1: raise RuntimeError('provider unavailable') - engines = SimpleNamespace(active_instances=instances, spawn_instance=spawn) - with pytest.raises(RuntimeError): - await claim_model(registry, SimpleNamespace(get=registry.read), engines, 'engine', 'a') - await claim_model(registry, SimpleNamespace(get=registry.read), engines, 'engine', 'a') - assert registry.values['slot:0'] == 'a' + registry = Registry({"routes": [{}, {"url": "latest"}], "trainer_demand:a": {}}) + + async def instances(_): + return [] + + engines = SimpleNamespace(active_instances=instances) + await claim_model( + registry, SimpleNamespace(get=registry.read), engines, "engine", "a" + ) + await claim_model( + registry, SimpleNamespace(get=registry.read), engines, "engine", "a" + ) + assert registry.values["slot:0"] == "a" + asyncio.run(check()) def test_stitch_run_switch_drains_old_requests_and_gates_new_identity(): async def check(): events = [] - async def reset(): events.append('reset') - async def pause(): events.append('pause') - async def resume(): events.append('resume') + + async def reset(): + events.append("reset") + + async def pause(): + events.append("pause") + + async def resume(): + events.append("resume") + engine = SimpleNamespace(reset=reset, pause=pause, resume=resume) - reconciler = AssignedReconciler(store=None, engine=engine, run_id='a') - reconciler.applied = VersionRef('a', 7) - token = _expected_run.set('a') + reconciler = AssignedReconciler(store=None, engine=engine, run_id="a") + reconciler.applied = VersionRef("a", 7) + token = _expected_run.set("a") async with reconciler.admit(VersionConstraint(min_version=7)): - switch = asyncio.create_task(reconciler._switch_run('b')) + switch = asyncio.create_task(reconciler._switch_run("b")) await asyncio.sleep(0) assert not switch.done() and events == [] await switch - assert events == ['pause', 'reset', 'resume'] - assert reconciler.applied == VersionRef('b', 0) + assert events == ["pause", "reset", "resume"] + assert reconciler.applied == VersionRef("b", 0) with pytest.raises(ConstraintUnmet): - async with reconciler.admit(VersionConstraint()): pass + async with reconciler.admit(VersionConstraint()): + pass _expected_run.reset(token) - token = _expected_run.set('b') + token = _expected_run.set("b") with pytest.raises(ConstraintUnmet): - async with reconciler.admit(VersionConstraint(min_version=1)): pass - reconciler.applied = VersionRef('b', 1) + async with reconciler.admit(VersionConstraint(min_version=1)): + pass + reconciler.applied = VersionRef("b", 1) async with reconciler.admit(VersionConstraint(min_version=1)) as served: - assert served == VersionRef('b', 1) + assert served == VersionRef("b", 1) _expected_run.reset(token) + asyncio.run(check()) def test_http_admission_rejects_old_handles_and_waits_for_run_switch(): async def check(): - registry = Registry({'slot:0': 'b'}) - engine = SimpleNamespace(base_url=lambda: 'http://unused', blocked_routes=lambda: ()) - reconciler = AssignedReconciler(store=None, engine=engine, run_id='a') - reconciler.applied = VersionRef('a', 999) + registry = Registry({"slot:0": "b"}) + engine = SimpleNamespace( + base_url=lambda: "http://unused", blocked_routes=lambda: () + ) + reconciler = AssignedReconciler(store=None, engine=engine, run_id="a") + reconciler.applied = VersionRef("a", 999) app = assigned_app(reconciler, engine, registry) - async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url='http://test') as client: - old = await client.post('/generate', json={'weight_run_id': 'a'}) + async with httpx.AsyncClient( + transport=httpx.ASGITransport(app=app), base_url="http://test" + ) as client: + old = await client.post("/generate", json={"weight_run_id": "a"}) assert old.status_code == 410 - new = await client.post('/generate', json={'weight_run_id': 'b', 'weight_version': {'min_version': 1}}) + new = await client.post( + "/generate", + json={"weight_run_id": "b", "weight_version": {"min_version": 1}}, + ) assert new.status_code == 409 # a/999 must never satisfy b/1. + asyncio.run(check()) def test_assigned_store_changes_pointer_without_rewriting_pinned_history(tmp_path): - from spindle.inference.full_bulletin import FFTSnapshotBulletin, PinnedFFTSnapshotStore + from spindle.inference.full_bulletin import ( + FFTSnapshotBulletin, + PinnedFFTSnapshotStore, + ) + board = FFTSnapshotBulletin(tmp_path) - board.claim('a') - values = {'slot:0': 'b'} - store = AssignedSnapshotStore(board, 'a', values) + board.claim("a") + values = {"slot:0": "b"} + store = AssignedSnapshotStore(board, "a", values) store.refresh() - assert store.read_pointer() == VersionRef('b', 0) - assert PinnedFFTSnapshotStore(board, VersionRef('a', 0)).read_pointer() == VersionRef('a', 0) + assert store.read_pointer() == VersionRef("b", 0) + assert PinnedFFTSnapshotStore( + board, VersionRef("a", 0) + ).read_pointer() == VersionRef("a", 0) def test_old_creation_retry_cannot_resurrect_a_placed_model(): async def check(): - registry = Registry({'slot:0': 'a', 'placement:a': {'engine': 'dead'}}) - async def instances(_): return [] - async def spawn(_): raise AssertionError('must not resurrect lost model') + registry = Registry({"slot:0": "a", "placement:a": {"engine": "dead"}}) + + async def instances(_): + return [] + + async def spawn(_): + raise AssertionError("must not resurrect lost model") + engines = SimpleNamespace(active_instances=instances, spawn_instance=spawn) - with pytest.raises(ValueError, match='trainer was lost'): - await claim_model(registry, SimpleNamespace(get=registry.read), engines, 'engine', 'a') + with pytest.raises(ValueError, match="trainer was lost"): + await claim_model( + registry, SimpleNamespace(get=registry.read), engines, "engine", "a" + ) + asyncio.run(check()) def test_cpu_run_switch_retires_after_drain_without_resetting_cache(monkeypatch): import spindle.inference.scoped_sidecar as module + async def check(): events = [] - async def reset(): raise AssertionError('CPU cache reset is unsupported') + + async def reset(): + raise AssertionError("CPU cache reset is unsupported") + async def retire(): - events.append('retired') - raise RuntimeError('simulated container termination') - monkeypatch.setattr(module, 'retire_replica', retire) - engine = SimpleNamespace(delta_update_mode='cpu', reset=reset) - reconciler = AssignedReconciler(store=None, engine=engine, run_id='a') - reconciler.applied = VersionRef('a', 1) + events.append("retired") + raise RuntimeError("simulated container termination") + + monkeypatch.setattr(module, "retire_replica", retire) + engine = SimpleNamespace(delta_update_mode="cpu", reset=reset) + reconciler = AssignedReconciler(store=None, engine=engine, run_id="a") + reconciler.applied = VersionRef("a", 1) async with reconciler.admit(VersionConstraint(min_version=1)): - switch = asyncio.create_task(reconciler._switch_run('b')) + switch = asyncio.create_task(reconciler._switch_run("b")) await asyncio.sleep(0) assert not switch.done() and events == [] - with pytest.raises(RuntimeError, match='simulated container termination'): + with pytest.raises(RuntimeError, match="simulated container termination"): await switch - assert events == ['retired'] - assert reconciler.applied == VersionRef('a', 1) # Never relabel old weights. + assert events == ["retired"] + assert reconciler.applied == VersionRef("a", 1) # Never relabel old weights. + asyncio.run(check()) diff --git a/tests/telemetry/test_metadata.py b/tests/telemetry/test_metadata.py index 065c135..35ce33e 100644 --- a/tests/telemetry/test_metadata.py +++ b/tests/telemetry/test_metadata.py @@ -60,18 +60,5 @@ async def run(): model_path=artifact.model_path, ) assert sampling.telemetry_tags == artifact.telemetry_tags - # Old records predate tagging. Adding labels must not change publication identity. - from spindle.control_plane.keys import sampler_artifact_key - - versioned_key = sampler_artifact_key( - plane._latest_sampler_model_path(creation.model.model_id, 7) - ) - old = await plane.kv.get(versioned_key) - old.pop("telemetry_tags") - await plane.kv.put(versioned_key, old) - rid = await plane.submit_sampler_export( - export_request(creation.model.model_id, seq_id=2, path="second") - ) - assert (await plane.retrieve(rid, timeout=1.0)).result is not None asyncio.run(run()) diff --git a/tests/test_deployments.py b/tests/test_deployments.py index a4fbfde..2e31d7f 100644 --- a/tests/test_deployments.py +++ b/tests/test_deployments.py @@ -15,7 +15,6 @@ from spindle.configs.qwen35_9b_lora_16k import Config as Parent from spindle.configuration import BaseConfig from spindle.control_plane import ControlPlane, create_control_plane_app -from spindle.control_plane.deployments import DeploymentRoutes from spindle.control_plane.keys import model_key from spindle.deployments import ( DeploymentConfig, @@ -101,37 +100,6 @@ def test_asset_paths_follow_the_model(): assert a.asset_path != b.asset_path -def test_routing_uses_deployment_order(): - small = resolved() - large = resolved(recipe("qwen35-9b-lora-64k")) - routes = DeploymentRoutes([small, large]) - assert routes.select(small.model, "lora").definition_id == small.definition_id - assert routes.capabilities()[0]["max_context_length"] == 16384 - validate_frontend([small.recipe, large.recipe]) - - routes = DeploymentRoutes([large]) - assert routes.select(small.model, "lora").definition_id == large.definition_id - assert routes.capabilities()[0]["max_context_length"] == 65536 - - other = resolved(recipe(name="other", model="org/other")) - routes = DeploymentRoutes([other]) - assert routes.select(other.model, "lora").definition_id == other.definition_id - assert {row["model_name"] for row in routes.capabilities()} == {other.model} - - -def test_sampling_uses_order_and_training_filters_parameterization(): - lora = resolved() - fft = resolved(recipe("qwen35-4b-fft-64k", model=lora.model)) - for first, second in ((lora, fft), (fft, lora)): - routes = DeploymentRoutes([first, second]) - assert routes.select(lora.model).definition_id == first.definition_id - assert routes.select(lora.model, "lora").definition_id == lora.definition_id - assert routes.select(lora.model, "full").definition_id == fft.definition_id - assert routes.select(second.definition_id).definition_id == second.definition_id - assert routes.select("missing") is None - assert routes.select(lora.definition_id, "full") is None - - def test_native_false_list_aliases_and_scalar_overrides(): parser = argparse.ArgumentParser() parser.add_argument("--use-feature", action="store_true")