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:
@@ -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
|
||||||
@@ -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
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user