Support scp GPU fleet sync
This commit is contained in:
@@ -182,6 +182,59 @@ def rsync_pull(config: FleetConfig, host: HostSpec, remote_path: str, local_path
|
|||||||
run_local(argv, cwd=config.project_root, capture_output=True, check=True)
|
run_local(argv, cwd=config.project_root, capture_output=True, check=True)
|
||||||
|
|
||||||
|
|
||||||
|
def scp_push(config: FleetConfig, host: HostSpec) -> None:
|
||||||
|
ensure_remote_dir(config, host, host.sync_remote_path)
|
||||||
|
local_src = str(config.sync.local_path.resolve()) + "/."
|
||||||
|
remote_dst = f"{host.ssh_alias}:{host.sync_remote_path.rstrip('/')}/"
|
||||||
|
argv = [
|
||||||
|
"scp",
|
||||||
|
"-o",
|
||||||
|
"BatchMode=yes",
|
||||||
|
"-o",
|
||||||
|
f"ConnectTimeout={config.ssh_timeout_sec}",
|
||||||
|
"-r",
|
||||||
|
"-p",
|
||||||
|
local_src,
|
||||||
|
remote_dst,
|
||||||
|
]
|
||||||
|
run_local(argv, cwd=config.project_root, capture_output=True, check=True)
|
||||||
|
|
||||||
|
|
||||||
|
def scp_pull(config: FleetConfig, host: HostSpec, remote_path: str, local_path: Path) -> None:
|
||||||
|
remote_src = remote_path
|
||||||
|
if remote_path.endswith("/"):
|
||||||
|
ensure_dir(local_path)
|
||||||
|
remote_src = remote_path.rstrip("/") + "/."
|
||||||
|
else:
|
||||||
|
ensure_dir(local_path.parent)
|
||||||
|
argv = [
|
||||||
|
"scp",
|
||||||
|
"-o",
|
||||||
|
"BatchMode=yes",
|
||||||
|
"-o",
|
||||||
|
f"ConnectTimeout={config.ssh_timeout_sec}",
|
||||||
|
"-r",
|
||||||
|
"-p",
|
||||||
|
f"{host.ssh_alias}:{remote_src}",
|
||||||
|
str(local_path),
|
||||||
|
]
|
||||||
|
run_local(argv, cwd=config.project_root, capture_output=True, check=True)
|
||||||
|
|
||||||
|
|
||||||
|
def sync_push(config: FleetConfig, host: HostSpec) -> None:
|
||||||
|
if config.sync.mode == "rsync":
|
||||||
|
rsync_push(config, host)
|
||||||
|
else:
|
||||||
|
scp_push(config, host)
|
||||||
|
|
||||||
|
|
||||||
|
def sync_pull(config: FleetConfig, host: HostSpec, remote_path: str, local_path: Path) -> None:
|
||||||
|
if config.sync.mode == "rsync":
|
||||||
|
rsync_pull(config, host, remote_path, local_path)
|
||||||
|
else:
|
||||||
|
scp_pull(config, host, remote_path, local_path)
|
||||||
|
|
||||||
|
|
||||||
def ensure_remote_dir(config: FleetConfig, host: HostSpec, remote_path: str) -> None:
|
def ensure_remote_dir(config: FleetConfig, host: HostSpec, remote_path: str) -> None:
|
||||||
run_ssh(config, host, f"mkdir -p {shlex.quote(remote_path)}", capture_output=True, check=True)
|
run_ssh(config, host, f"mkdir -p {shlex.quote(remote_path)}", capture_output=True, check=True)
|
||||||
|
|
||||||
@@ -206,8 +259,10 @@ def load_config(path: Path) -> FleetConfig:
|
|||||||
local_path=relative_to_root(project_root, sync_raw.get("local_path"), project_root),
|
local_path=relative_to_root(project_root, sync_raw.get("local_path"), project_root),
|
||||||
exclude=[str(item) for item in sync_raw.get("exclude", [])],
|
exclude=[str(item) for item in sync_raw.get("exclude", [])],
|
||||||
)
|
)
|
||||||
if sync.mode != "rsync":
|
if sync.mode not in {"rsync", "scp"}:
|
||||||
raise FleetError(f"unsupported sync.mode: {sync.mode}")
|
raise FleetError(f"unsupported sync.mode: {sync.mode}")
|
||||||
|
if sync.mode == "scp" and sync.exclude:
|
||||||
|
raise FleetError("sync.exclude is not supported for sync.mode=scp")
|
||||||
|
|
||||||
scheduler_raw = raw.get("scheduler", {})
|
scheduler_raw = raw.get("scheduler", {})
|
||||||
scheduler = SchedulerSpec(
|
scheduler = SchedulerSpec(
|
||||||
@@ -639,7 +694,7 @@ def harvest_run(config: FleetConfig, manifest: dict[str, Any]) -> dict[str, Any]
|
|||||||
|
|
||||||
local_base = ensure_dir(config.artifacts_dir / refreshed["run_id"])
|
local_base = ensure_dir(config.artifacts_dir / refreshed["run_id"])
|
||||||
remote_run_dir = refreshed["remote_run_dir"].rstrip("/")
|
remote_run_dir = refreshed["remote_run_dir"].rstrip("/")
|
||||||
rsync_pull(config, host, f"{remote_run_dir}/", local_base / "remote_run")
|
sync_pull(config, host, f"{remote_run_dir}/", local_base / "remote_run")
|
||||||
|
|
||||||
for artifact in refreshed.get("artifacts", []):
|
for artifact in refreshed.get("artifacts", []):
|
||||||
artifact_remote = f"{refreshed['remote_sync_path'].rstrip('/')}/{artifact}"
|
artifact_remote = f"{refreshed['remote_sync_path'].rstrip('/')}/{artifact}"
|
||||||
@@ -652,7 +707,7 @@ def harvest_run(config: FleetConfig, manifest: dict[str, Any]) -> dict[str, Any]
|
|||||||
check=False,
|
check=False,
|
||||||
)
|
)
|
||||||
if check.returncode == 0:
|
if check.returncode == 0:
|
||||||
rsync_pull(config, host, artifact_remote, target)
|
sync_pull(config, host, artifact_remote, target)
|
||||||
|
|
||||||
refreshed["harvested_at"] = utc_now()
|
refreshed["harvested_at"] = utc_now()
|
||||||
write_run_manifest(config, refreshed)
|
write_run_manifest(config, refreshed)
|
||||||
@@ -866,7 +921,7 @@ def dispatch_jobs(
|
|||||||
)
|
)
|
||||||
continue
|
continue
|
||||||
if host.name not in synced_hosts:
|
if host.name not in synced_hosts:
|
||||||
rsync_push(config, host)
|
sync_push(config, host)
|
||||||
synced_hosts.add(host.name)
|
synced_hosts.add(host.name)
|
||||||
manifest = launch_job(config, host, job, gpu_ids)
|
manifest = launch_job(config, host, job, gpu_ids)
|
||||||
manifests.append(manifest)
|
manifests.append(manifest)
|
||||||
|
|||||||
Reference in New Issue
Block a user