From 071dfaf42cd033d9b16a9e457664b3fc78d246c6 Mon Sep 17 00:00:00 2001 From: dela Date: Wed, 26 Aug 2026 10:00:06 +0800 Subject: [PATCH] Sanitize SwanLab env before login so 0.9 nested project does not crash OpenBayes sets SWANLAB_PROJECT as a string; swanlab>=0.9 parses that as ProjectSettings and raises QuoteAwareEnvSettingsSource. Drop it, keep SWANLAB_PROJ_NAME, and share run-id extraction with train_k3. --- kda/training/swanlab_env.py | 41 +++++++++++++++++++++++++++ tests/integration/test_swanlab_env.py | 39 +++++++++++++++++++++++++ train_k3.py | 30 ++++++-------------- 3 files changed, 89 insertions(+), 21 deletions(-) create mode 100644 kda/training/swanlab_env.py create mode 100644 tests/integration/test_swanlab_env.py diff --git a/kda/training/swanlab_env.py b/kda/training/swanlab_env.py new file mode 100644 index 0000000..eb60dc3 --- /dev/null +++ b/kda/training/swanlab_env.py @@ -0,0 +1,41 @@ +"""Sanitize SwanLab env before import/init. + +swanlab>=0.9 ``Settings.project`` is a nested model. A string +``SWANLAB_PROJECT`` (OpenBayes and older docs) makes pydantic raise +``error parsing value for field "project" from source +_QuoteAwareEnvSettingsSource``. Project name belongs in +``SWANLAB_PROJ_NAME`` / ``init(project=...)``. +""" + +from __future__ import annotations + +import os + + +def prepare_swanlab_env(default_project: str = "kda") -> str: + """Drop nested ``SWANLAB_PROJECT``, keep a plain project name. + + Must run before ``import swanlab`` / ``swanlab.login`` / ``init``. + """ + raw = os.environ.pop("SWANLAB_PROJECT", None) + name = os.environ.get("SWANLAB_PROJ_NAME") or raw or default_project + name = str(name).strip().strip("\"'") + if not name or name[0] in "{[": + name = default_project + os.environ.pop("SWANLAB_PROJECT", None) + os.environ["SWANLAB_PROJ_NAME"] = name + return name + + +def swanlab_run_id(run) -> str | None: + for attr in ("id", "run_id"): + val = getattr(run, attr, None) + if isinstance(val, str) and val: + return val + public = getattr(run, "public", None) + if public is not None: + for attr in ("cloud_run_id", "run_id", "id"): + val = getattr(public, attr, None) + if isinstance(val, str) and val: + return val + return None diff --git a/tests/integration/test_swanlab_env.py b/tests/integration/test_swanlab_env.py new file mode 100644 index 0000000..d64f232 --- /dev/null +++ b/tests/integration/test_swanlab_env.py @@ -0,0 +1,39 @@ +import os + +from kda.training.swanlab_env import prepare_swanlab_env, swanlab_run_id + + +def test_prepare_drops_string_project(monkeypatch): + monkeypatch.setenv("SWANLAB_PROJECT", "kda") + monkeypatch.delenv("SWANLAB_PROJ_NAME", raising=False) + assert prepare_swanlab_env() == "kda" + assert "SWANLAB_PROJECT" not in os.environ + assert os.environ["SWANLAB_PROJ_NAME"] == "kda" + + +def test_prepare_strips_quotes(monkeypatch): + monkeypatch.setenv("SWANLAB_PROJECT", '"kda"') + monkeypatch.delenv("SWANLAB_PROJ_NAME", raising=False) + assert prepare_swanlab_env() == "kda" + assert "SWANLAB_PROJECT" not in os.environ + + +def test_prepare_prefers_proj_name(monkeypatch): + monkeypatch.setenv("SWANLAB_PROJECT", "ignored") + monkeypatch.setenv("SWANLAB_PROJ_NAME", "mine") + assert prepare_swanlab_env() == "mine" + assert os.environ["SWANLAB_PROJ_NAME"] == "mine" + + +def test_prepare_rejects_json_blob(monkeypatch): + monkeypatch.setenv("SWANLAB_PROJECT", '{"name": "x"}') + monkeypatch.delenv("SWANLAB_PROJ_NAME", raising=False) + assert prepare_swanlab_env() == "kda" + + +def test_swanlab_run_id(): + class _Run: + id = "ilgne5ro" + + assert swanlab_run_id(_Run()) == "ilgne5ro" + assert swanlab_run_id(object()) is None diff --git a/train_k3.py b/train_k3.py index ccfe310..8530289 100644 --- a/train_k3.py +++ b/train_k3.py @@ -22,6 +22,7 @@ from kda.models.causal_lm import CausalLM from kda.models.k3_config import K3Config from kda.training.data import iter_indexed, load_pretrain_chunks, load_tokenizer from kda.training.schedule import lr_scale, tokens_per_micro, total_opt_steps +from kda.training.swanlab_env import prepare_swanlab_env, swanlab_run_id from kda.training.toy import load_ckpt _TOY_TRAIN = { @@ -66,25 +67,14 @@ def _set_lr(optim: torch.optim.Optimizer, lr: float) -> None: group["lr"] = lr -def _swanlab_run_id(run) -> str | None: - for attr in ("id", "run_id"): - val = getattr(run, attr, None) - if isinstance(val, str) and val: - return val - public = getattr(run, "public", None) - if public is not None: - for attr in ("cloud_run_id", "run_id", "id"): - val = getattr(public, attr, None) - if isinstance(val, str) and val: - return val - return None - - -def _init_swanlab(cfg: K3Config, args: argparse.Namespace, resume_id: str | None = None): +def _init_swanlab( + cfg: K3Config, args: argparse.Namespace, resume_id: str | None = None +): """Cloud monitor if SWANLAB_API_KEY is set; otherwise no-op.""" key = os.environ.get("SWANLAB_API_KEY") if not key: return None + project = prepare_swanlab_env() try: import swanlab except ImportError: @@ -92,8 +82,6 @@ def _init_swanlab(cfg: K3Config, args: argparse.Namespace, resume_id: str | None return None try: swanlab.login(api_key=key, save=False) - # swanlab 0.9 Settings.project is nested; a string SWANLAB_PROJECT env crashes init. - project = os.environ.pop("SWANLAB_PROJECT", None) or "kda" run_id = resume_id or os.environ.get("SWANLAB_RUN_ID") init_kw = dict( project=project, @@ -125,7 +113,7 @@ def _init_swanlab(cfg: K3Config, args: argparse.Namespace, resume_id: str | None init_kw["resume"] = True print(f"swanlab resume id={run_id}") run = swanlab.init(**init_kw) - got = _swanlab_run_id(run) + got = swanlab_run_id(run) if got: args.swanlab_id = got print(f"swanlab run id {got}") @@ -523,9 +511,9 @@ def main() -> None: if elapsed > 0: metrics["train/tok_s"] = (tokens - tokens_at_t0) / elapsed - ended = ( - args.max_tokens is not None and tokens >= args.max_tokens - ) or (args.max_tokens is None and micro_step >= args.steps) + ended = (args.max_tokens is not None and tokens >= args.max_tokens) or ( + args.max_tokens is None and micro_step >= args.steps + ) log_now = micro_step % args.log_every == 0 or micro_step == 1 or ended eval_now = micro_step % args.eval_every == 0 or micro_step == 1 or ended ckpt_now = micro_step % args.ckpt_every == 0 or ended