Skip to content

Commit 482f0f1

Browse files
committed
Add Tinker checkpoint delete endpoint
1 parent 152ebec commit 482f0f1

3 files changed

Lines changed: 257 additions & 5 deletions

File tree

.claude/docs/tinker.md

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -37,6 +37,7 @@ All endpoints are under `/api/v1/`. Requests are async -- submit via POST, get a
3737
| `/save_weights` | POST | Save full training checkpoint (weights + optimizer) |
3838
| `/save_weights_for_sampler` | POST | Sync weights to inference engines |
3939
| `/load_weights` | POST | Load a previously saved checkpoint |
40+
| `/training_runs/{unique_id}/checkpoints/{checkpoint_id}` | DELETE | Delete a saved checkpoint archive from disk |
4041
| `/retrieve_future` | POST | Long-poll for async result (300s timeout) |
4142
| `/healthz` | GET | Liveness check |
4243

@@ -52,6 +53,7 @@ All endpoints are under `/api/v1/`. Requests are async -- submit via POST, get a
5253
- **Persistent**: `save_weights_for_sampler(name="...")` -- syncs to inference engines AND writes HF checkpoint to disk. Expensive.
5354
- **Ephemeral**: `save_weights_and_get_sampling_client(name="...")` -- syncs to inference engines only, skips disk write. Triggered when `sampling_session_seq_id` is present in the request.
5455
- In RL loops, always prefer ephemeral mode; reserve persistent saves for periodic checkpoints.
56+
- Delete persistent checkpoints with `DELETE /training_runs/{unique_id}/checkpoints/{checkpoint_id}` after they are no longer needed. This removes the saved archive from `checkpoints_base`; it does not unload the live model.
5557

5658
## Testing
5759

@@ -74,4 +76,4 @@ uv run --extra dev --extra jax pytest tests/tinker/test_api.py -v
7476

7577
- **Token shifting**: Tinker pre-shifts inputs/targets; SkyRL-Train shifts internally. The backend appends the last target token to reconstruct full sequences during batch conversion -- be careful if modifying `prepare_model_pass_batch`.
7678
- **Left-padding**: SkyRL-Train expects left-padded tensors. The backend handles this during batch prep.
77-
- **API models vs internal types**: `api.py` defines its own Pydantic models (e.g., `api.ForwardBackwardInput`) that mirror but differ from `types.ForwardBackwardInput`. Each API model has a `.to_types()` method for conversion. Do not confuse the two.
79+
- **API models vs internal types**: `api.py` defines its own Pydantic models (e.g., `api.ForwardBackwardInput`) that mirror but differ from `types.ForwardBackwardInput`. Each API model has a `.to_types()` method for conversion. Do not confuse the two.

skyrl/tinker/api.py

Lines changed: 96 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,8 @@
11
import asyncio
22
import os
33
import random
4+
import re
5+
import shutil
46
import signal
57
import threading
68
import time
@@ -1219,8 +1221,59 @@ async def validate_checkpoint(
12191221
if checkpoint_db.status == CheckpointStatus.FAILED:
12201222
raise HTTPException(status_code=500, detail=f"Checkpoint creation failed: {checkpoint_db.error_message}")
12211223

1222-
subdir = "sampler_weights" if checkpoint_type == types.CheckpointType.SAMPLER else ""
1223-
return request.app.state.engine_config.checkpoints_base / unique_id / subdir / f"{checkpoint_id}.tar.gz"
1224+
return checkpoint_file_path(request, unique_id, checkpoint_id, checkpoint_type)
1225+
1226+
1227+
def checkpoint_file_path(
1228+
request: Request, unique_id: str, checkpoint_id: str, checkpoint_type: types.CheckpointType
1229+
) -> Any:
1230+
checkpoint_dir = request.app.state.engine_config.checkpoints_base / unique_id
1231+
if checkpoint_type == types.CheckpointType.SAMPLER:
1232+
checkpoint_dir = checkpoint_dir / "sampler_weights"
1233+
return checkpoint_dir / f"{checkpoint_id}.tar.gz"
1234+
1235+
1236+
def parse_checkpoint_delete_path(
1237+
checkpoint_path: str, checkpoint_type: types.CheckpointType | None
1238+
) -> tuple[str, types.CheckpointType | None]:
1239+
path_kind, separator, checkpoint_id = checkpoint_path.partition("/")
1240+
if separator:
1241+
if path_kind == "weights":
1242+
inferred_checkpoint_type = types.CheckpointType.TRAINING
1243+
elif path_kind == "sampler_weights":
1244+
inferred_checkpoint_type = types.CheckpointType.SAMPLER
1245+
else:
1246+
raise HTTPException(status_code=400, detail=f"Invalid checkpoint path: {checkpoint_path}")
1247+
1248+
if checkpoint_type is not None and checkpoint_type != inferred_checkpoint_type:
1249+
raise HTTPException(status_code=400, detail="checkpoint_type does not match checkpoint path")
1250+
else:
1251+
checkpoint_id = checkpoint_path
1252+
inferred_checkpoint_type = checkpoint_type
1253+
1254+
if not re.fullmatch(ID_PATTERN, checkpoint_id) or len(checkpoint_id) > ID_MAX_LENGTH:
1255+
raise HTTPException(status_code=422, detail="Invalid checkpoint_id")
1256+
1257+
return checkpoint_id, inferred_checkpoint_type
1258+
1259+
1260+
def delete_checkpoint_file(checkpoint_path: Any) -> None:
1261+
if checkpoint_path.is_dir():
1262+
if hasattr(checkpoint_path, "rmtree"):
1263+
checkpoint_path.rmtree()
1264+
else:
1265+
shutil.rmtree(checkpoint_path)
1266+
else:
1267+
try:
1268+
checkpoint_path.unlink()
1269+
except FileNotFoundError:
1270+
return
1271+
1272+
with suppress(OSError):
1273+
checkpoint_path.parent.rmdir()
1274+
if checkpoint_path.parent.name == "sampler_weights":
1275+
with suppress(OSError):
1276+
checkpoint_path.parent.parent.rmdir()
12241277

12251278

12261279
@app.get("/api/v1/training_runs")
@@ -1303,6 +1356,38 @@ async def download_checkpoint_archive(
13031356
return StreamingResponse(file_buffer, media_type="application/octet-stream", headers=headers)
13041357

13051358

1359+
@app.delete("/api/v1/training_runs/{unique_id}/checkpoints/{checkpoint_path:path}", status_code=204)
1360+
async def delete_checkpoint(
1361+
request: Request,
1362+
unique_id: str = fastapi.Path(..., pattern=ID_PATTERN, max_length=ID_MAX_LENGTH),
1363+
checkpoint_path: str = fastapi.Path(...),
1364+
checkpoint_type: types.CheckpointType | None = fastapi.Query(None),
1365+
session: AsyncSession = Depends(get_session),
1366+
) -> None:
1367+
"""Delete a saved checkpoint artifact and its database row."""
1368+
checkpoint_id, resolved_checkpoint_type = parse_checkpoint_delete_path(checkpoint_path, checkpoint_type)
1369+
checkpoint_db = None
1370+
if resolved_checkpoint_type is None:
1371+
resolved_checkpoint_type = types.CheckpointType.TRAINING
1372+
checkpoint_db = await session.get(CheckpointDB, (unique_id, checkpoint_id, resolved_checkpoint_type))
1373+
if checkpoint_db is None:
1374+
resolved_checkpoint_type = types.CheckpointType.SAMPLER
1375+
checkpoint_db = await session.get(CheckpointDB, (unique_id, checkpoint_id, resolved_checkpoint_type))
1376+
else:
1377+
checkpoint_db = await session.get(CheckpointDB, (unique_id, checkpoint_id, resolved_checkpoint_type))
1378+
1379+
if not checkpoint_db:
1380+
raise HTTPException(status_code=404, detail=f"Checkpoint not found: {unique_id}/{checkpoint_id}")
1381+
1382+
if checkpoint_db.status == CheckpointStatus.PENDING:
1383+
raise HTTPException(status_code=425, detail="Checkpoint is still being created")
1384+
1385+
path = checkpoint_file_path(request, unique_id, checkpoint_id, resolved_checkpoint_type)
1386+
await asyncio.to_thread(delete_checkpoint_file, path)
1387+
await session.delete(checkpoint_db)
1388+
await session.commit()
1389+
1390+
13061391
@app.get("/api/v1/training_runs/{unique_id}/checkpoints")
13071392
async def list_checkpoints(
13081393
unique_id: str = fastapi.Path(..., pattern=ID_PATTERN, max_length=ID_MAX_LENGTH),
@@ -1379,12 +1464,19 @@ async def root():
13791464
"name": "Tinker API Mock",
13801465
"version": "0.0.1",
13811466
"endpoints": {
1382-
"models": ["/api/v1/create_model", "/api/v1/get_info", "/api/v1/training_runs/{model_id}"],
1467+
"models": [
1468+
"/api/v1/create_model",
1469+
"/api/v1/get_info",
1470+
"/api/v1/training_runs/{model_id}",
1471+
],
13831472
"training": ["/api/v1/forward_backward", "/api/v1/optim_step"],
13841473
"futures": ["/api/v1/retrieve_future"],
13851474
"service": ["/api/v1/get_server_capabilities"],
13861475
"telemetry": ["/api/v1/telemetry"],
1387-
"checkpoints": ["/api/v1/training_runs/{unique_id}/checkpoints"],
1476+
"checkpoints": [
1477+
"/api/v1/training_runs/{unique_id}/checkpoints",
1478+
"DELETE /api/v1/training_runs/{unique_id}/checkpoints/{checkpoint_id}",
1479+
],
13881480
"download": [
13891481
"/api/v1/training_runs/{unique_id}/checkpoints/{checkpoint_id}/archive",
13901482
"/api/v1/training_runs/{unique_id}/checkpoints/{checkpoint_id}/download",
Lines changed: 158 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,158 @@
1+
"""Fast tests for checkpoint deletion API behavior."""
2+
3+
import asyncio
4+
from collections.abc import AsyncGenerator, Iterator
5+
from datetime import datetime, timezone
6+
from pathlib import Path
7+
8+
import pytest
9+
from fastapi.testclient import TestClient
10+
from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine
11+
from sqlmodel import SQLModel
12+
from sqlmodel.ext.asyncio.session import AsyncSession
13+
14+
from skyrl.tinker import types
15+
from skyrl.tinker.api import app, get_session
16+
from skyrl.tinker.config import EngineConfig
17+
from skyrl.tinker.db_models import CheckpointDB, CheckpointStatus, ModelDB, SessionDB
18+
19+
MODEL_ID = "model_fast_delete"
20+
21+
22+
@pytest.fixture
23+
def checkpoint_api(tmp_path: Path) -> Iterator[tuple[TestClient, Path, AsyncEngine]]:
24+
db_path = tmp_path / "checkpoint_delete.db"
25+
checkpoint_base = tmp_path / "checkpoints"
26+
db_engine = create_async_engine(f"sqlite+aiosqlite:///{db_path}")
27+
28+
async def setup_db() -> None:
29+
async with db_engine.begin() as conn:
30+
await conn.run_sync(SQLModel.metadata.create_all)
31+
32+
async def override_get_session() -> AsyncGenerator[AsyncSession, None]:
33+
async with AsyncSession(db_engine) as session:
34+
yield session
35+
36+
asyncio.run(setup_db())
37+
previous_state = app.state._state.copy()
38+
app.state.db_engine = db_engine
39+
app.state.engine_config = EngineConfig(
40+
base_model="test-model",
41+
checkpoints_base=checkpoint_base,
42+
database_url=f"sqlite:///{db_path}",
43+
)
44+
app.dependency_overrides[get_session] = override_get_session
45+
client = TestClient(app)
46+
47+
try:
48+
yield client, checkpoint_base, db_engine
49+
finally:
50+
client.close()
51+
app.dependency_overrides.pop(get_session, None)
52+
app.state._state.clear()
53+
app.state._state.update(previous_state)
54+
asyncio.run(db_engine.dispose())
55+
56+
57+
async def seed_model_and_checkpoints(
58+
db_engine: AsyncEngine,
59+
checkpoint_ids: list[str],
60+
checkpoint_type: types.CheckpointType = types.CheckpointType.TRAINING,
61+
) -> None:
62+
async with AsyncSession(db_engine) as session:
63+
if await session.get(SessionDB, "session_fast_delete") is None:
64+
session.add(SessionDB(session_id="session_fast_delete", tags=[], sdk_version="test"))
65+
if await session.get(ModelDB, MODEL_ID) is None:
66+
session.add(
67+
ModelDB(
68+
model_id=MODEL_ID,
69+
base_model="test-model",
70+
lora_config={"rank": 1},
71+
status="created",
72+
request_id=1,
73+
session_id="session_fast_delete",
74+
)
75+
)
76+
for checkpoint_id in checkpoint_ids:
77+
session.add(
78+
CheckpointDB(
79+
model_id=MODEL_ID,
80+
checkpoint_id=checkpoint_id,
81+
checkpoint_type=checkpoint_type,
82+
status=CheckpointStatus.COMPLETED,
83+
completed_at=datetime.now(timezone.utc),
84+
)
85+
)
86+
await session.commit()
87+
88+
89+
def write_training_checkpoint(checkpoint_base: Path, checkpoint_id: str, directory: bool = False) -> Path:
90+
checkpoint_path = checkpoint_base / MODEL_ID / f"{checkpoint_id}.tar.gz"
91+
checkpoint_path.parent.mkdir(parents=True, exist_ok=True)
92+
if directory:
93+
checkpoint_path.mkdir()
94+
(checkpoint_path / "checkpoint_0").write_text("tiny")
95+
else:
96+
checkpoint_path.write_text("tiny")
97+
return checkpoint_path
98+
99+
100+
def write_sampler_checkpoint(checkpoint_base: Path, checkpoint_id: str) -> Path:
101+
checkpoint_path = checkpoint_base / MODEL_ID / "sampler_weights" / f"{checkpoint_id}.tar.gz"
102+
checkpoint_path.parent.mkdir(parents=True, exist_ok=True)
103+
checkpoint_path.write_text("tiny")
104+
return checkpoint_path
105+
106+
107+
def listed_checkpoint_ids(client: TestClient) -> set[str]:
108+
response = client.get(f"/api/v1/training_runs/{MODEL_ID}/checkpoints")
109+
assert response.status_code == 200
110+
return {checkpoint["checkpoint_id"] for checkpoint in response.json()["checkpoints"]}
111+
112+
113+
def test_delete_checkpoint_removes_saved_artifact_and_list_entry(
114+
checkpoint_api: tuple[TestClient, Path, AsyncEngine],
115+
) -> None:
116+
client, checkpoint_base, db_engine = checkpoint_api
117+
asyncio.run(seed_model_and_checkpoints(db_engine, ["delete_training"]))
118+
training_checkpoint = write_training_checkpoint(checkpoint_base, "delete_training", directory=True)
119+
120+
response = client.delete(f"/api/v1/training_runs/{MODEL_ID}/checkpoints/weights/delete_training")
121+
122+
assert response.status_code == 204
123+
assert not training_checkpoint.exists()
124+
assert "delete_training" not in listed_checkpoint_ids(client)
125+
126+
asyncio.run(
127+
seed_model_and_checkpoints(db_engine, ["delete_sampler"], checkpoint_type=types.CheckpointType.SAMPLER)
128+
)
129+
sampler_checkpoint = write_sampler_checkpoint(checkpoint_base, "delete_sampler")
130+
131+
response = client.delete(f"/api/v1/training_runs/{MODEL_ID}/checkpoints/delete_sampler")
132+
133+
assert response.status_code == 204
134+
assert not sampler_checkpoint.exists()
135+
assert "delete_sampler" not in listed_checkpoint_ids(client)
136+
137+
138+
def test_delete_even_checkpoints_leaves_odd_checkpoints_listed(
139+
checkpoint_api: tuple[TestClient, Path, AsyncEngine],
140+
) -> None:
141+
client, checkpoint_base, db_engine = checkpoint_api
142+
checkpoint_ids = ["1", "2", "3", "4", "5"]
143+
asyncio.run(seed_model_and_checkpoints(db_engine, checkpoint_ids))
144+
checkpoint_files = {
145+
checkpoint_id: write_training_checkpoint(checkpoint_base, checkpoint_id) for checkpoint_id in checkpoint_ids
146+
}
147+
148+
assert listed_checkpoint_ids(client) == {"1", "2", "3", "4", "5"}
149+
150+
assert client.delete(f"/api/v1/training_runs/{MODEL_ID}/checkpoints/2").status_code == 204
151+
assert client.delete(f"/api/v1/training_runs/{MODEL_ID}/checkpoints/4").status_code == 204
152+
153+
assert checkpoint_files["1"].exists()
154+
assert not checkpoint_files["2"].exists()
155+
assert checkpoint_files["3"].exists()
156+
assert not checkpoint_files["4"].exists()
157+
assert checkpoint_files["5"].exists()
158+
assert listed_checkpoint_ids(client) == {"1", "3", "5"}

0 commit comments

Comments
 (0)