Skip to content

Ray Launcher

The experiment launcher is graphids.exp.ray_backend. It accepts a typed RunConfig, constructs Ray TorchTrainer directly, and runs the worker loop that builds data/model/trainer objects from YAML config.

Current stages:

  • fit / test -> Lightning trainer launch with config-driven data/model instantiation, Ray Lightning strategy/environment, and Ray checkpoint/metric reporting.

graphids.exp.ray_backend

ray_backend

Ray Train launcher for GraphIDS experiment YAML.

RunSummary dataclass

RunSummary(run_dir: str, status: str, stage: str, name: str, last_event: str | None = None, error: str | None = None, extra: dict[str, Any] = dict())

Minimal status summary for UI/readout code.

build_component

build_component(spec: Any, **build_kwargs: Any) -> Any

Resolve a spec and call build() when the resolved object supports it.

Source code in graphids/exp/ray_backend.py
def build_component(spec: Any, **build_kwargs: Any) -> Any:
    """Resolve a spec and call ``build()`` when the resolved object supports it."""
    resolved = resolve_spec(spec)
    builder = getattr(resolved, "build", None)
    if callable(builder):
        if build_kwargs:
            params = signature(builder).parameters
            filtered = {k: v for k, v in build_kwargs.items() if k in params}
            try:
                return builder(**filtered)
            except TypeError:
                pass
        return builder()
    return resolved

launch_run

launch_run(run: RunConfig, *, address: str | None = None) -> RunSummary

Run a GraphIDS experiment through Ray Train.

Source code in graphids/exp/ray_backend.py
def launch_run(run: RunConfig, *, address: str | None = None) -> RunSummary:
    """Run a GraphIDS experiment through Ray Train."""
    try:
        import ray
        from ray.train import CheckpointConfig, ScalingConfig
        from ray.train import RunConfig as RayRunConfig
        from ray.train.torch import TorchTrainer
    except ModuleNotFoundError as exc:
        raise RuntimeError("Ray is not installed. Install the project dependencies before launching experiments.") from exc

    if not ray.is_initialized():
        ray.init(address=address, ignore_reinit_error=True)

    resources = run.resources
    devices = run.payload.trainer.get("devices")
    num_workers = devices if isinstance(devices, int) and devices > 1 else 1
    use_gpu = resources.accelerator == "gpu" or float(resources.gpus_per_worker) > 0
    resources_per_worker: dict[str, float | int] = {"CPU": max(1, int(resources.cpus_per_worker))}
    if float(resources.gpus_per_worker) > 0:
        resources_per_worker["GPU"] = float(resources.gpus_per_worker)

    run_dir = Path(run.outputs.run_dir)
    ray_run_kwargs: dict[str, Any] = {"storage_path": str(run_dir.parent), "name": run_dir.name}
    callbacks = run.payload.callbacks or {}
    callback_values = callbacks.values() if isinstance(callbacks, Mapping) else callbacks
    for callback in callback_values:
        if not isinstance(callback, Mapping):
            continue
        callback_type = str(callback.get("type") or callback.get("class_path") or "")
        if "checkpoint" not in callback_type.lower():
            continue
        checkpoint_kwargs: dict[str, Any] = {}
        if isinstance(callback.get("save_top_k"), int) and callback["save_top_k"] > 0:
            checkpoint_kwargs["num_to_keep"] = callback["save_top_k"]
        if isinstance(callback.get("monitor"), str):
            checkpoint_kwargs["checkpoint_score_attribute"] = callback["monitor"]
        if callback.get("mode") in {"min", "max"}:
            checkpoint_kwargs["checkpoint_score_order"] = callback["mode"]
        if checkpoint_kwargs:
            ray_run_kwargs["checkpoint_config"] = CheckpointConfig(**checkpoint_kwargs)
        break

    result = TorchTrainer(
        _worker_loop,
        train_loop_config={"run": run.model_dump(mode="json")},
        scaling_config=ScalingConfig(
            num_workers=num_workers,
            use_gpu=use_gpu,
            resources_per_worker=resources_per_worker,
        ),
        run_config=RayRunConfig(**ray_run_kwargs),
    ).fit()

    append_event(
        run.outputs.run_dir,
        EventRecord(
            status="finished",
            stage=run.stage,
            message="ray_result",
            details={"path": str(getattr(result, "path", "") or ""), "metrics": _jsonish(getattr(result, "metrics", {}) or {})},
        ),
        name=run.outputs.events_name,
    )
    summary = summarize_run(run.outputs.run_dir)
    if summary is None:
        raise RuntimeError(f"Ray run finished without a GraphIDS manifest: {run.outputs.run_dir}")
    return summary

probe_ray_train_imports

probe_ray_train_imports() -> dict[str, str]

Import the Ray APIs GraphIDS relies on and return their module paths.

Source code in graphids/exp/ray_backend.py
def probe_ray_train_imports() -> dict[str, str]:
    """Import the Ray APIs GraphIDS relies on and return their module paths."""
    try:
        from ray.train import CheckpointConfig, ScalingConfig
        from ray.train import RunConfig as RayRunConfig
        from ray.train.lightning import (
            RayDDPStrategy,
            RayLightningEnvironment,
            RayTrainReportCallback,
            prepare_trainer,
        )
        from ray.train.torch import TorchTrainer
    except ModuleNotFoundError as exc:
        raise RuntimeError("Ray is not installed. Install the project dependencies before launching experiments.") from exc

    return {
        "RunConfig": RayRunConfig.__module__,
        "ScalingConfig": ScalingConfig.__module__,
        "CheckpointConfig": CheckpointConfig.__module__,
        "TorchTrainer": TorchTrainer.__module__,
        "prepare_trainer": prepare_trainer.__module__,
        "RayDDPStrategy": RayDDPStrategy.__module__,
        "RayLightningEnvironment": RayLightningEnvironment.__module__,
        "RayTrainReportCallback": RayTrainReportCallback.__module__,
    }

resolve_spec

resolve_spec(spec: Any) -> Any

Resolve a nested primitive/class-path spec without calling build().

Source code in graphids/exp/ray_backend.py
def resolve_spec(spec: Any) -> Any:
    """Resolve a nested primitive/class-path spec without calling ``build()``."""
    if hasattr(spec, "model_dump") and not isinstance(spec, Mapping):
        spec = spec.model_dump(mode="json")
    if isinstance(spec, list):
        return [resolve_spec(item) for item in spec]
    if not isinstance(spec, Mapping):
        return spec

    if "class_path" in spec:
        class_path = str(spec["class_path"])
        module_path, _, class_name = class_path.rpartition(".")
        cls = getattr(importlib.import_module(module_path), class_name)
        init_args = {k: resolve_spec(v) for k, v in dict(spec.get("init_args") or {}).items()}
        return cls(**init_args)

    if "type" in spec:
        from graphids import primitives as primitive_mod

        factory = getattr(primitive_mod, str(spec["type"]), None)
        if callable(factory):
            kwargs = {k: resolve_spec(v) for k, v in spec.items() if k != "type"}
            return factory(**kwargs)

    return {k: resolve_spec(v) for k, v in spec.items()}