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.
This commit is contained in:
dela
2026-08-26 10:00:06 +08:00
parent 9a4862a866
commit 071dfaf42c
3 changed files with 89 additions and 21 deletions
+41
View File
@@ -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
+39
View File
@@ -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
+9 -21
View File
@@ -22,6 +22,7 @@ from kda.models.causal_lm import CausalLM
from kda.models.k3_config import K3Config from kda.models.k3_config import K3Config
from kda.training.data import iter_indexed, load_pretrain_chunks, load_tokenizer 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.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 from kda.training.toy import load_ckpt
_TOY_TRAIN = { _TOY_TRAIN = {
@@ -66,25 +67,14 @@ def _set_lr(optim: torch.optim.Optimizer, lr: float) -> None:
group["lr"] = lr group["lr"] = lr
def _swanlab_run_id(run) -> str | None: def _init_swanlab(
for attr in ("id", "run_id"): cfg: K3Config, args: argparse.Namespace, resume_id: str | None = None
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):
"""Cloud monitor if SWANLAB_API_KEY is set; otherwise no-op.""" """Cloud monitor if SWANLAB_API_KEY is set; otherwise no-op."""
key = os.environ.get("SWANLAB_API_KEY") key = os.environ.get("SWANLAB_API_KEY")
if not key: if not key:
return None return None
project = prepare_swanlab_env()
try: try:
import swanlab import swanlab
except ImportError: except ImportError:
@@ -92,8 +82,6 @@ def _init_swanlab(cfg: K3Config, args: argparse.Namespace, resume_id: str | None
return None return None
try: try:
swanlab.login(api_key=key, save=False) 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") run_id = resume_id or os.environ.get("SWANLAB_RUN_ID")
init_kw = dict( init_kw = dict(
project=project, project=project,
@@ -125,7 +113,7 @@ def _init_swanlab(cfg: K3Config, args: argparse.Namespace, resume_id: str | None
init_kw["resume"] = True init_kw["resume"] = True
print(f"swanlab resume id={run_id}") print(f"swanlab resume id={run_id}")
run = swanlab.init(**init_kw) run = swanlab.init(**init_kw)
got = _swanlab_run_id(run) got = swanlab_run_id(run)
if got: if got:
args.swanlab_id = got args.swanlab_id = got
print(f"swanlab run id {got}") print(f"swanlab run id {got}")
@@ -523,9 +511,9 @@ def main() -> None:
if elapsed > 0: if elapsed > 0:
metrics["train/tok_s"] = (tokens - tokens_at_t0) / elapsed metrics["train/tok_s"] = (tokens - tokens_at_t0) / elapsed
ended = ( ended = (args.max_tokens is not None and tokens >= args.max_tokens) or (
args.max_tokens is not None and tokens >= args.max_tokens args.max_tokens is None and micro_step >= args.steps
) 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 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 eval_now = micro_step % args.eval_every == 0 or micro_step == 1 or ended
ckpt_now = micro_step % args.ckpt_every == 0 or ended ckpt_now = micro_step % args.ckpt_every == 0 or ended