feat: publish downloadable wakeword trainer source
This commit is contained in:
@@ -0,0 +1,603 @@
|
||||
"""Standard-library host launcher for the downloadable JRich Wakeword Trainer."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import hashlib
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import shutil
|
||||
import subprocess
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
RUNNER_VERSION = "0.2.1-jrich-2"
|
||||
DEFAULT_IMAGE = (
|
||||
"git.jimandkrista.com/jr-public/jrich-wakeword-training"
|
||||
"@sha256:cef5b1e769a6160bbde69f67df76f0b865f02c27794a3da096398c18cc9a3c2f"
|
||||
)
|
||||
DEFAULT_CACHE = Path.home() / ".cache" / "jrich-wakeword" / "livekit-0.2.1"
|
||||
DEFAULT_WORK_ROOT = Path.cwd() / "jrich-wakeword-runs"
|
||||
MIN_FREE_BYTES = 40 * 1024**3
|
||||
PRESET_PATH = (
|
||||
Path(__file__).resolve().parents[1]
|
||||
/ "deploy"
|
||||
/ "wakeword-training"
|
||||
/ "presets"
|
||||
/ "manual-poc-v1.json"
|
||||
)
|
||||
_SLUG = re.compile(r"^[a-z0-9]+(?:-[a-z0-9]+)*$")
|
||||
_MODEL_NAME = re.compile(r"^[A-Za-z0-9][A-Za-z0-9_-]*$")
|
||||
|
||||
|
||||
def _canonical(value: object) -> str:
|
||||
return json.dumps(value, ensure_ascii=False, sort_keys=True, separators=(",", ":"))
|
||||
|
||||
|
||||
def _sha256(value: bytes) -> str:
|
||||
return hashlib.sha256(value).hexdigest()
|
||||
|
||||
|
||||
def _deep_merge(base: dict[str, Any], override: dict[str, Any]) -> dict[str, Any]:
|
||||
result = dict(base)
|
||||
for key, value in override.items():
|
||||
if key not in base:
|
||||
raise ValueError(f"unknown pipeline override: {key}")
|
||||
if isinstance(base[key], dict) and isinstance(value, dict):
|
||||
result[key] = _deep_merge(base[key], value)
|
||||
else:
|
||||
result[key] = value
|
||||
return result
|
||||
|
||||
|
||||
def default_pipeline_config(
|
||||
*,
|
||||
model_name: str,
|
||||
family_slug: str,
|
||||
version: int,
|
||||
phrases: tuple[str, ...],
|
||||
near_misses: tuple[str, ...],
|
||||
) -> dict[str, Any]:
|
||||
"""Return the measured, versioned manual-POC profile before user overrides."""
|
||||
|
||||
preset = json.loads(PRESET_PATH.read_text(encoding="utf-8"))
|
||||
if not isinstance(preset, dict) or preset.get("preset_version") != "manual-poc-v1":
|
||||
raise ValueError("the bundled manual POC preset is malformed")
|
||||
config = preset.get("pipeline_config")
|
||||
if not isinstance(config, dict):
|
||||
raise ValueError("the bundled manual POC preset has no pipeline_config")
|
||||
return {
|
||||
"model_name": model_name,
|
||||
"family_slug": family_slug,
|
||||
"version": version,
|
||||
"target_phrases": list(phrases),
|
||||
"custom_negative_phrases": list(near_misses),
|
||||
**config,
|
||||
}
|
||||
|
||||
|
||||
def build_setup_command(
|
||||
*, image_ref: str, gpu_uuid: str, cache_dir: Path, config_path: Path
|
||||
) -> tuple[str, ...]:
|
||||
"""Build the explicit networked cache-population command; training stays offline."""
|
||||
|
||||
cache_dir.mkdir(parents=True, exist_ok=True)
|
||||
return (
|
||||
"docker",
|
||||
"run",
|
||||
"--rm",
|
||||
"--init",
|
||||
"--gpus",
|
||||
f"device={gpu_uuid}",
|
||||
"--read-only",
|
||||
"--cap-drop",
|
||||
"ALL",
|
||||
"--security-opt",
|
||||
"no-new-privileges",
|
||||
"--mount",
|
||||
f"type=bind,src={cache_dir.resolve()},dst=/data",
|
||||
"--mount",
|
||||
f"type=bind,src={config_path.resolve()},dst=/config.json,readonly",
|
||||
"--tmpfs",
|
||||
"/tmp:rw,noexec,nosuid,size=4g",
|
||||
image_ref,
|
||||
"livekit-wakeword",
|
||||
"setup",
|
||||
"--config",
|
||||
"/config.json",
|
||||
)
|
||||
|
||||
|
||||
def _run(
|
||||
command: tuple[str, ...] | list[str], *, capture: bool = False
|
||||
) -> subprocess.CompletedProcess:
|
||||
return subprocess.run(command, check=False, capture_output=capture, text=True)
|
||||
|
||||
|
||||
def _inspect_image(image_ref: str) -> dict[str, Any] | None:
|
||||
completed = _run(
|
||||
["docker", "image", "inspect", "--format", "{{json .}}", image_ref],
|
||||
capture=True,
|
||||
)
|
||||
if completed.returncode != 0:
|
||||
return None
|
||||
try:
|
||||
value = json.loads(completed.stdout)
|
||||
except json.JSONDecodeError as exc:
|
||||
raise ValueError("Docker returned malformed image metadata") from exc
|
||||
if not isinstance(value, dict):
|
||||
raise ValueError("Docker returned malformed image metadata")
|
||||
return value
|
||||
|
||||
|
||||
def _image_identity(image_ref: str) -> str:
|
||||
if image_ref.startswith("sha256:") and len(image_ref) == 71:
|
||||
return image_ref.lower()
|
||||
|
||||
inspected = _inspect_image(image_ref)
|
||||
if inspected is None:
|
||||
print(f"Downloading pinned training runtime: {image_ref}", file=sys.stderr)
|
||||
pulled = _run(["docker", "pull", image_ref])
|
||||
if pulled.returncode != 0:
|
||||
raise ValueError(f"training image download failed: {image_ref}")
|
||||
inspected = _inspect_image(image_ref)
|
||||
if inspected is None:
|
||||
raise ValueError(
|
||||
f"training image is unavailable after download: {image_ref}"
|
||||
)
|
||||
|
||||
if "@sha256:" in image_ref:
|
||||
expected = image_ref.lower()
|
||||
repo_digests = inspected.get("RepoDigests", [])
|
||||
if not isinstance(repo_digests, list) or expected not in {
|
||||
str(value).lower() for value in repo_digests
|
||||
}:
|
||||
raise ValueError("Docker did not retain the requested registry manifest digest")
|
||||
|
||||
digest = str(inspected.get("Id", "")).lower()
|
||||
if not digest.startswith("sha256:") or len(digest) != 71:
|
||||
raise ValueError("Docker returned an invalid immutable image ID")
|
||||
return digest
|
||||
|
||||
|
||||
def _default_gpu_uuid() -> str:
|
||||
completed = _run(
|
||||
["nvidia-smi", "--query-gpu=uuid", "--format=csv,noheader"], capture=True
|
||||
)
|
||||
values = [line.strip() for line in completed.stdout.splitlines() if line.strip()]
|
||||
if completed.returncode != 0 or not values:
|
||||
raise ValueError("no NVIDIA GPU UUID was found; verify the driver and nvidia-smi")
|
||||
if len(values) > 1:
|
||||
raise ValueError("multiple NVIDIA GPUs found; choose one with --gpu-uuid")
|
||||
return values[0]
|
||||
|
||||
|
||||
def _prompt(label: str, default: str | None = None) -> str:
|
||||
suffix = f" [{default}]" if default else ""
|
||||
value = input(f"{label}{suffix}: ").strip()
|
||||
return value or (default or "")
|
||||
|
||||
|
||||
def _slug(value: str) -> str:
|
||||
candidate = re.sub(r"[^a-z0-9]+", "-", value.lower()).strip("-")
|
||||
if not _SLUG.fullmatch(candidate):
|
||||
raise ValueError("family slug must use lowercase letters, numbers, and single hyphens")
|
||||
return candidate
|
||||
|
||||
|
||||
def _interactive_values(args: argparse.Namespace) -> None:
|
||||
if args.phrase:
|
||||
return
|
||||
print("JRich Wakeword Trainer — the defaults below match the measured RTX 3090 POC.")
|
||||
primary = _prompt("Primary spoken phrase", "Hey Jay Rich")
|
||||
variants = _prompt("Additional spoken phrases, comma-separated", "Yo Jay Rich")
|
||||
near_misses = _prompt("Near-miss phrases that must not activate, comma-separated", "Jay Rich")
|
||||
label = _prompt("Model name", "jrich")
|
||||
args.phrase = [primary, *[item.strip() for item in variants.split(",") if item.strip()]]
|
||||
args.near_miss = [item.strip() for item in near_misses.split(",") if item.strip()]
|
||||
args.model_name = args.model_name or label
|
||||
args.family_slug = args.family_slug or _slug(label)
|
||||
|
||||
|
||||
def _load_override(path: Path | None) -> dict[str, Any]:
|
||||
if path is None:
|
||||
return {}
|
||||
value = json.loads(path.read_text(encoding="utf-8"))
|
||||
if not isinstance(value, dict):
|
||||
raise ValueError("pipeline override must be a JSON object")
|
||||
return value
|
||||
|
||||
|
||||
def _file_sha256(path: Path) -> str:
|
||||
digest = hashlib.sha256()
|
||||
with path.open("rb") as file:
|
||||
for block in iter(lambda: file.read(1024 * 1024), b""):
|
||||
digest.update(block)
|
||||
return digest.hexdigest()
|
||||
|
||||
|
||||
def write_local_source_lock(work_root: Path, cache_dir: Path) -> tuple[Path, str]:
|
||||
"""Write a path-independent, content-addressed inventory of the local cache."""
|
||||
|
||||
request_root = work_root / "requests"
|
||||
request_root.mkdir(parents=True, exist_ok=True)
|
||||
cache_dir.mkdir(parents=True, exist_ok=True)
|
||||
files: list[dict[str, object]] = []
|
||||
for path in sorted(cache_dir.rglob("*")):
|
||||
if path.is_symlink():
|
||||
raise ValueError(f"training cache must not contain symlinks: {path}")
|
||||
if not path.is_file():
|
||||
continue
|
||||
files.append(
|
||||
{
|
||||
"relative_path": path.relative_to(cache_dir).as_posix(),
|
||||
"size_bytes": path.stat().st_size,
|
||||
"sha256": _file_sha256(path),
|
||||
}
|
||||
)
|
||||
value = {
|
||||
"schema_version": 1,
|
||||
"run_kind": "local_experiment",
|
||||
"sources": [
|
||||
{
|
||||
"source_id": "livekit-setup-cache-v0.2.1",
|
||||
"kind": "user-local-cache",
|
||||
"publication_approved": False,
|
||||
"files": files,
|
||||
}
|
||||
],
|
||||
}
|
||||
encoded = (_canonical(value) + "\n").encode("utf-8")
|
||||
digest = _sha256(encoded)
|
||||
path = request_root / f"sources-{digest}.lock.json"
|
||||
if not path.exists():
|
||||
path.write_bytes(encoded)
|
||||
path.chmod(0o444)
|
||||
return path, digest
|
||||
|
||||
|
||||
def resolve_plan(args: argparse.Namespace) -> dict[str, Any]:
|
||||
phrases = tuple(" ".join(value.split()) for value in (args.phrase or ()) if value.strip())
|
||||
if not phrases:
|
||||
raise ValueError("at least one --phrase is required in non-interactive mode")
|
||||
near_misses = tuple(
|
||||
" ".join(value.split()) for value in (args.near_miss or ()) if value.strip()
|
||||
)
|
||||
model_name = (args.model_name or _slug(phrases[0]).replace("-", "_")).strip()
|
||||
if not _MODEL_NAME.fullmatch(model_name):
|
||||
raise ValueError(
|
||||
"model name must start with a letter or number and contain only letters, "
|
||||
"numbers, underscores, or hyphens"
|
||||
)
|
||||
if args.version < 1:
|
||||
raise ValueError("version must be a positive integer")
|
||||
family_slug = _slug(args.family_slug or model_name)
|
||||
image_digest = _image_identity(args.image_ref)
|
||||
gpu_uuid = args.gpu_uuid or _default_gpu_uuid()
|
||||
cache_dir = args.cache_dir.resolve()
|
||||
work_root = args.work_root.resolve()
|
||||
pipeline = default_pipeline_config(
|
||||
model_name=model_name,
|
||||
family_slug=family_slug,
|
||||
version=args.version,
|
||||
phrases=phrases,
|
||||
near_misses=near_misses,
|
||||
)
|
||||
pipeline = _deep_merge(pipeline, _load_override(args.pipeline_config))
|
||||
identity = {
|
||||
"model_name": model_name,
|
||||
"family_slug": family_slug,
|
||||
"version": args.version,
|
||||
"target_phrases": list(phrases),
|
||||
"custom_negative_phrases": list(near_misses),
|
||||
}
|
||||
divergent = [name for name, value in identity.items() if pipeline.get(name) != value]
|
||||
if divergent:
|
||||
raise ValueError(
|
||||
"pipeline override cannot change request identity fields: " + ", ".join(divergent)
|
||||
)
|
||||
source_lock, source_digest = write_local_source_lock(work_root, cache_dir)
|
||||
request = {
|
||||
"schema_version": 1,
|
||||
"primary_phrase": phrases[0],
|
||||
"variants": list(phrases[1:]),
|
||||
"near_misses": list(near_misses),
|
||||
"evaluation_profile": "local-balanced-v1",
|
||||
"source_lock_path": str(source_lock),
|
||||
"source_lock_sha256": source_digest,
|
||||
"intended_release_id": f"{family_slug}-v{args.version}",
|
||||
"runner_image_digest": image_digest,
|
||||
"pipeline_config": pipeline,
|
||||
"run_kind": "local_experiment",
|
||||
"derived_from_run_id": args.reuse_from,
|
||||
}
|
||||
# Keep the local wizard and future API worker on the exact same fail-closed
|
||||
# validation boundary. Import shape differs only between the extracted
|
||||
# archive (sibling module) and repository tests (scripts package).
|
||||
try:
|
||||
from wakeword_training import RunKind, WakewordTrainingRequest
|
||||
except ModuleNotFoundError:
|
||||
from scripts.wakeword_training import RunKind, WakewordTrainingRequest
|
||||
|
||||
WakewordTrainingRequest(
|
||||
primary_phrase=request["primary_phrase"],
|
||||
variants=tuple(request["variants"]),
|
||||
near_misses=tuple(request["near_misses"]),
|
||||
evaluation_profile=request["evaluation_profile"],
|
||||
source_lock_path=Path(request["source_lock_path"]),
|
||||
source_lock_sha256=request["source_lock_sha256"],
|
||||
intended_release_id=request["intended_release_id"],
|
||||
runner_image_digest=request["runner_image_digest"],
|
||||
pipeline_config=request["pipeline_config"],
|
||||
run_kind=RunKind.LOCAL_EXPERIMENT,
|
||||
derived_from_run_id=request["derived_from_run_id"],
|
||||
).canonical_payload()
|
||||
return {
|
||||
"schema_version": 1,
|
||||
"runner_version": RUNNER_VERSION,
|
||||
"request": request,
|
||||
"runtime": {
|
||||
"image_ref": image_digest,
|
||||
"image_source": args.image_ref,
|
||||
"gpu_uuid": gpu_uuid,
|
||||
"cache_dir": str(cache_dir),
|
||||
"work_root": str(work_root),
|
||||
"hardware_profile": "shared-conservative"
|
||||
if args.allow_active_gpu_processes
|
||||
else "exclusive",
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _preflight(plan: dict[str, Any]) -> None:
|
||||
runtime = plan["runtime"]
|
||||
for executable in ("docker", "nvidia-smi"):
|
||||
if shutil.which(executable) is None:
|
||||
raise ValueError(f"required host command is missing: {executable}")
|
||||
docker = _run(["docker", "info"], capture=True)
|
||||
if docker.returncode != 0:
|
||||
raise ValueError("Docker daemon is unavailable to this user")
|
||||
gpu = _run(
|
||||
[
|
||||
"nvidia-smi",
|
||||
f"--id={runtime['gpu_uuid']}",
|
||||
"--query-gpu=uuid,name,memory.total,memory.free,driver_version",
|
||||
"--format=csv,noheader,nounits",
|
||||
],
|
||||
capture=True,
|
||||
)
|
||||
if gpu.returncode != 0:
|
||||
raise ValueError("the selected GPU UUID is unavailable")
|
||||
container_gpu = _run(
|
||||
[
|
||||
"docker",
|
||||
"run",
|
||||
"--rm",
|
||||
"--gpus",
|
||||
f"device={runtime['gpu_uuid']}",
|
||||
"--network",
|
||||
"none",
|
||||
"--read-only",
|
||||
"--cap-drop",
|
||||
"ALL",
|
||||
"--security-opt",
|
||||
"no-new-privileges",
|
||||
runtime["image_ref"],
|
||||
"python",
|
||||
"-c",
|
||||
"import torch; assert torch.cuda.is_available(); print(torch.cuda.get_device_name())",
|
||||
],
|
||||
capture=True,
|
||||
)
|
||||
if container_gpu.returncode != 0:
|
||||
raise ValueError("NVIDIA Container Toolkit did not expose the selected GPU to the image")
|
||||
for key in ("cache_dir", "work_root"):
|
||||
path = Path(runtime[key])
|
||||
path.mkdir(parents=True, exist_ok=True)
|
||||
free = shutil.disk_usage(path).free
|
||||
if free < MIN_FREE_BYTES:
|
||||
raise ValueError(f"{path} has less than the recommended 40 GiB free")
|
||||
print(f"GPU: {gpu.stdout.strip()}", file=sys.stderr)
|
||||
print(f"Container CUDA: {container_gpu.stdout.strip()}", file=sys.stderr)
|
||||
print(f"Host CPUs: {os.cpu_count() or 'unknown'}", file=sys.stderr)
|
||||
print("Preflight passed; training will run with network disabled.", file=sys.stderr)
|
||||
|
||||
|
||||
def _write_request(plan: dict[str, Any]) -> Path:
|
||||
work_root = Path(plan["runtime"]["work_root"])
|
||||
encoded = (_canonical(plan["request"]) + "\n").encode("utf-8")
|
||||
path = work_root / "requests" / f"request-{_sha256(encoded)}.json"
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
path.write_bytes(encoded)
|
||||
return path
|
||||
|
||||
|
||||
def _setup(plan: dict[str, Any]) -> None:
|
||||
runtime = plan["runtime"]
|
||||
config = dict(plan["request"]["pipeline_config"])
|
||||
config["data_dir"] = "/data"
|
||||
config["output_dir"] = "/tmp/setup-output"
|
||||
setup_path = Path(runtime["work_root"]) / "requests" / "setup-config.json"
|
||||
setup_path.write_text(_canonical(config) + "\n", encoding="utf-8")
|
||||
command = build_setup_command(
|
||||
image_ref=runtime["image_ref"],
|
||||
gpu_uuid=runtime["gpu_uuid"],
|
||||
cache_dir=Path(runtime["cache_dir"]),
|
||||
config_path=setup_path,
|
||||
)
|
||||
completed = _run(command)
|
||||
if completed.returncode != 0:
|
||||
raise ValueError("LiveKit asset setup failed")
|
||||
|
||||
|
||||
def _train(plan: dict[str, Any], request_path: Path, args: argparse.Namespace) -> int:
|
||||
runtime = plan["runtime"]
|
||||
command = [
|
||||
sys.executable,
|
||||
str(Path(__file__).with_name("run-wakeword-training")),
|
||||
"--request",
|
||||
str(request_path),
|
||||
"--work-root",
|
||||
runtime["work_root"],
|
||||
"--enable",
|
||||
"--maintenance-window-active",
|
||||
"--image-ref",
|
||||
runtime["image_ref"],
|
||||
"--gpu-uuid",
|
||||
runtime["gpu_uuid"],
|
||||
"--approved-inputs-dir",
|
||||
runtime["cache_dir"],
|
||||
"--timeout-s",
|
||||
str(args.timeout_s),
|
||||
"--stream-output",
|
||||
]
|
||||
if args.allow_active_gpu_processes:
|
||||
command.append("--allow-active-gpu-processes")
|
||||
if args.retry_of:
|
||||
command.extend(("--retry-of", args.retry_of))
|
||||
return _run(command).returncode
|
||||
|
||||
|
||||
def _find_run(work_root: Path, run_id: str) -> tuple[Path, dict[str, Any]]:
|
||||
if len(run_id) != 32 or any(character not in "0123456789abcdef" for character in run_id):
|
||||
raise ValueError("run ID must be 32 lowercase hexadecimal characters")
|
||||
matches: list[tuple[Path, dict[str, Any]]] = []
|
||||
for run_dir in (work_root / "runs").glob(f"*-{run_id}"):
|
||||
state_path = run_dir / "state.json"
|
||||
if not state_path.is_file():
|
||||
continue
|
||||
state = json.loads(state_path.read_text(encoding="utf-8"))
|
||||
if isinstance(state, dict) and state.get("run_id") == run_id:
|
||||
matches.append((run_dir, state))
|
||||
if len(matches) != 1:
|
||||
raise ValueError("run ID was not found under the selected work root")
|
||||
return matches[0]
|
||||
|
||||
|
||||
def run_local_operation(args: argparse.Namespace) -> int | None:
|
||||
operation = next(
|
||||
(
|
||||
(name, run_id)
|
||||
for name, run_id in (
|
||||
("status", args.status),
|
||||
("logs", args.logs),
|
||||
("cancel", args.cancel),
|
||||
)
|
||||
if run_id is not None
|
||||
),
|
||||
None,
|
||||
)
|
||||
if operation is None:
|
||||
return None
|
||||
name, run_id = operation
|
||||
run_dir, state = _find_run(args.work_root.resolve(), run_id)
|
||||
if name == "status":
|
||||
result_path = run_dir / "training-result.json"
|
||||
result = (
|
||||
json.loads(result_path.read_text(encoding="utf-8"))
|
||||
if result_path.is_file()
|
||||
else None
|
||||
)
|
||||
print(json.dumps({"run_dir": str(run_dir), "state": state, "result": result}, indent=2))
|
||||
return 0
|
||||
if name == "logs":
|
||||
for label in ("stdout.log", "stderr.log"):
|
||||
print(f"== {label} ==")
|
||||
path = run_dir / label
|
||||
if path.is_file():
|
||||
print(path.read_text(encoding="utf-8", errors="replace"), end="")
|
||||
return 0
|
||||
if state.get("status") in {
|
||||
"completed",
|
||||
"pending_publication",
|
||||
"failed",
|
||||
"canceled",
|
||||
"timed_out",
|
||||
"abandoned",
|
||||
}:
|
||||
raise ValueError("terminal run cannot be canceled")
|
||||
(run_dir / "cancel-requested").write_text("requested\n", encoding="utf-8")
|
||||
print(f"Cancellation requested for {run_id}; the runner will stop its own container.")
|
||||
return 0
|
||||
|
||||
|
||||
def _parser() -> argparse.ArgumentParser:
|
||||
parser = argparse.ArgumentParser(description="Train a LiveKit Wakeword model in pinned Docker")
|
||||
parser.add_argument("--non-interactive", action="store_true")
|
||||
parser.add_argument("--plan", action="store_true", help="print resolved JSON without running")
|
||||
operation = parser.add_mutually_exclusive_group()
|
||||
operation.add_argument("--status", metavar="RUN_ID", help="show one local attempt")
|
||||
operation.add_argument("--logs", metavar="RUN_ID", help="show one local attempt's logs")
|
||||
operation.add_argument("--cancel", metavar="RUN_ID", help="request safe cancellation")
|
||||
parser.add_argument("--phrase", action="append", help="spoken phrase; repeat for alternatives")
|
||||
parser.add_argument("--near-miss", action="append", help="phrase that must not activate")
|
||||
parser.add_argument("--model-name")
|
||||
parser.add_argument("--family-slug")
|
||||
parser.add_argument("--version", type=int, default=1)
|
||||
parser.add_argument("--pipeline-config", type=Path, help="JSON deep overrides")
|
||||
parser.add_argument("--image-ref", default=os.getenv("JR_WAKEWORD_IMAGE", DEFAULT_IMAGE))
|
||||
parser.add_argument("--gpu-uuid")
|
||||
parser.add_argument("--cache-dir", type=Path, default=DEFAULT_CACHE)
|
||||
parser.add_argument("--work-root", type=Path, default=DEFAULT_WORK_ROOT)
|
||||
parser.add_argument("--skip-setup", action="store_true")
|
||||
parser.add_argument(
|
||||
"--setup-only",
|
||||
action="store_true",
|
||||
help="populate the pinned input cache, print its source-lock digest, and exit",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--reuse-from",
|
||||
help="terminal local run ID whose sealed generated corpus seeds a new attempt",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--retry-of", help="matching terminal run ID to retry as a new immutable attempt"
|
||||
)
|
||||
parser.add_argument("--allow-active-gpu-processes", action="store_true")
|
||||
parser.add_argument("--timeout-s", type=float, default=6 * 60 * 60)
|
||||
return parser
|
||||
|
||||
|
||||
def main(argv: list[str] | None = None) -> int:
|
||||
args = _parser().parse_args(argv)
|
||||
try:
|
||||
if getattr(args, "setup_only", False) and args.skip_setup:
|
||||
raise ValueError("--setup-only and --skip-setup cannot be combined")
|
||||
operation_result = run_local_operation(args)
|
||||
if operation_result is not None:
|
||||
return operation_result
|
||||
if not args.non_interactive:
|
||||
_interactive_values(args)
|
||||
plan = resolve_plan(args)
|
||||
if args.plan:
|
||||
print(json.dumps(plan, indent=2, sort_keys=True))
|
||||
return 0
|
||||
print(json.dumps(plan, indent=2, sort_keys=True))
|
||||
if not args.non_interactive and _prompt("Continue with this configuration", "yes") != "yes":
|
||||
return 0
|
||||
_preflight(plan)
|
||||
if not args.skip_setup:
|
||||
_setup(plan)
|
||||
refreshed = resolve_plan(args)
|
||||
if refreshed["request"]["source_lock_sha256"] != plan["request"][
|
||||
"source_lock_sha256"
|
||||
]:
|
||||
plan = refreshed
|
||||
print("Final content-locked request after setup:", file=sys.stderr)
|
||||
print(json.dumps(plan, indent=2, sort_keys=True))
|
||||
if getattr(args, "setup_only", False):
|
||||
print(
|
||||
"WAKEWORD_TRAINING_SOURCE_LOCK_SHA256="
|
||||
+ plan["request"]["source_lock_sha256"]
|
||||
)
|
||||
return 0
|
||||
request_path = _write_request(plan)
|
||||
return _train(plan, request_path, args)
|
||||
except (OSError, ValueError, json.JSONDecodeError) as exc:
|
||||
print(f"error: {exc}", file=sys.stderr)
|
||||
return 2
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
Reference in New Issue
Block a user