430917d8f2
Adds `giant train-submit`/`giant rollout-submit`, mirroring `giant analyze submit`'s CPU-job pattern but for single remote-GPU jobs: +RemoteJob/ RequestGPUs, TARGET.ProvidesEtpCeph instead of the local-only ProvidesETPResources, and a self-contained condor/ run dir (wrapper, submit description, and a CondorJobMeta sidecar recording what was submitted and the assigned cluster id) so a run stays traceable after the fact. Training jobs re-check for last.pt on every wrapper invocation so a preempted job resumes instead of restarting. Also adds `giant new-run` to scaffold a run's config.toml + run dir (with collision-free naming via the new shared `default_out_dir`) ahead of submission, and factors router-flag parsing into `_router_cli_overrides` so `train` and `new-run` resolve it identically. Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
170 lines
5.1 KiB
Python
170 lines
5.1 KiB
Python
import uuid
|
|
|
|
import yaml
|
|
|
|
from giant.cli import (
|
|
_CEPH_PREDICTIONS,
|
|
_resolve_prediction_output,
|
|
_write_prediction_ref,
|
|
)
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# _resolve_prediction_output
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
def test_non_ceph_path_goes_to_data_parent(tmp_path):
|
|
data = tmp_path / "pools" / "pbwo4" / "full.manifest"
|
|
out, dataset_path, pred_uuid = _resolve_prediction_output(data, None)
|
|
|
|
assert out.parent == data.parent
|
|
assert out.name == f"{pred_uuid}.parquet"
|
|
assert dataset_path == data.resolve()
|
|
|
|
|
|
def test_ceph_path_goes_to_central_store(tmp_path, monkeypatch):
|
|
# Patch resolve() so /ceph/... exists on any machine running the tests.
|
|
ceph_data = _CEPH_PREDICTIONS.parent / "pools" / "pbwo4" / "full.manifest"
|
|
monkeypatch.setattr(
|
|
"giant.cli.Path.resolve",
|
|
lambda self: ceph_data if self == ceph_data else self.absolute(),
|
|
)
|
|
out, _, pred_uuid = _resolve_prediction_output(ceph_data, None)
|
|
|
|
assert out.parent == _CEPH_PREDICTIONS
|
|
assert out.name == f"{pred_uuid}.parquet"
|
|
|
|
|
|
def test_explicit_out_is_used_as_is(tmp_path):
|
|
data = tmp_path / "data.parquet"
|
|
explicit = tmp_path / "my_output.parquet"
|
|
out, _, _ = _resolve_prediction_output(data, explicit)
|
|
|
|
assert out == explicit
|
|
|
|
|
|
def test_uuid_is_valid(tmp_path):
|
|
data = tmp_path / "data.parquet"
|
|
_, _, pred_uuid = _resolve_prediction_output(data, None)
|
|
parsed = uuid.UUID(pred_uuid)
|
|
assert parsed.version == 4
|
|
|
|
|
|
def test_each_call_produces_a_distinct_uuid(tmp_path):
|
|
data = tmp_path / "data.parquet"
|
|
_, _, uuid1 = _resolve_prediction_output(data, None)
|
|
_, _, uuid2 = _resolve_prediction_output(data, None)
|
|
assert uuid1 != uuid2
|
|
|
|
|
|
def test_pinned_pred_uuid_is_used_as_is(tmp_path):
|
|
# `rollout-submit` pins the uuid at submit time so its condor run-dir,
|
|
# the output filename, and the eventual YAML sidecar all agree.
|
|
data = tmp_path / "data.parquet"
|
|
pinned = "abcd1234-abcd-4abc-9abc-abcdabcdabcd"
|
|
out, _, pred_uuid = _resolve_prediction_output(data, None, pred_uuid=pinned)
|
|
|
|
assert pred_uuid == pinned
|
|
assert out.name == f"{pinned}.parquet"
|
|
|
|
|
|
def test_dataset_path_is_resolved(tmp_path):
|
|
data = tmp_path / "data.parquet"
|
|
_, dataset_path, _ = _resolve_prediction_output(data, None)
|
|
assert dataset_path.is_absolute()
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# _write_prediction_ref
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
def test_ref_file_created_in_checkpoint_dir(tmp_path):
|
|
ckpt_dir = tmp_path / "checkpoints" / "run1"
|
|
ckpt_dir.mkdir(parents=True)
|
|
checkpoint = ckpt_dir / "best.pt"
|
|
checkpoint.touch()
|
|
|
|
out = tmp_path / "predictions" / "abc.parquet"
|
|
dataset = tmp_path / "pools" / "pbwo4" / "full.manifest"
|
|
pred_uuid = str(uuid.uuid4())
|
|
|
|
ref_path = _write_prediction_ref(checkpoint, pred_uuid, out, dataset)
|
|
|
|
assert ref_path == ckpt_dir / f"{pred_uuid}.yaml"
|
|
assert ref_path.exists()
|
|
|
|
|
|
def test_ref_yaml_contains_expected_fields(tmp_path):
|
|
ckpt_dir = tmp_path / "checkpoints"
|
|
ckpt_dir.mkdir()
|
|
checkpoint = ckpt_dir / "best.pt"
|
|
checkpoint.touch()
|
|
|
|
out = tmp_path / "pred.parquet"
|
|
dataset = tmp_path / "full.manifest"
|
|
pred_uuid = str(uuid.uuid4())
|
|
|
|
ref_path = _write_prediction_ref(checkpoint, pred_uuid, out, dataset)
|
|
data = yaml.safe_load(ref_path.read_text())
|
|
|
|
assert data["prediction_id"] == pred_uuid
|
|
assert data["output"] == str(out)
|
|
assert data["dataset"] == str(dataset)
|
|
assert data["checkpoint"] == str(checkpoint.resolve())
|
|
assert "timestamp" in data
|
|
assert "comment" not in data
|
|
|
|
|
|
def test_ref_yaml_includes_comment_when_provided(tmp_path):
|
|
ckpt_dir = tmp_path / "checkpoints"
|
|
ckpt_dir.mkdir()
|
|
checkpoint = ckpt_dir / "best.pt"
|
|
checkpoint.touch()
|
|
|
|
out = tmp_path / "pred.parquet"
|
|
dataset = tmp_path / "full.manifest"
|
|
pred_uuid = str(uuid.uuid4())
|
|
|
|
ref_path = _write_prediction_ref(
|
|
checkpoint, pred_uuid, out, dataset, comment="baseline sweep run 3"
|
|
)
|
|
data = yaml.safe_load(ref_path.read_text())
|
|
|
|
assert data["comment"] == "baseline sweep run 3"
|
|
|
|
|
|
def test_ref_timestamp_is_iso_format(tmp_path):
|
|
from datetime import datetime
|
|
|
|
ckpt_dir = tmp_path / "checkpoints"
|
|
ckpt_dir.mkdir()
|
|
checkpoint = ckpt_dir / "best.pt"
|
|
checkpoint.touch()
|
|
|
|
pred_uuid = str(uuid.uuid4())
|
|
ref_path = _write_prediction_ref(
|
|
checkpoint, pred_uuid, tmp_path / "p.parquet", tmp_path / "d"
|
|
)
|
|
data = yaml.safe_load(ref_path.read_text())
|
|
|
|
# Must parse without error and be timezone-aware (UTC).
|
|
ts = datetime.fromisoformat(data["timestamp"])
|
|
assert ts.tzinfo is not None
|
|
|
|
|
|
def test_ref_checkpoint_path_is_absolute(tmp_path):
|
|
ckpt_dir = tmp_path / "checkpoints"
|
|
ckpt_dir.mkdir()
|
|
checkpoint = ckpt_dir / "best.pt"
|
|
checkpoint.touch()
|
|
|
|
pred_uuid = str(uuid.uuid4())
|
|
ref_path = _write_prediction_ref(
|
|
checkpoint, pred_uuid, tmp_path / "p.parquet", tmp_path / "d"
|
|
)
|
|
data = yaml.safe_load(ref_path.read_text())
|
|
|
|
assert data["checkpoint"].startswith("/")
|