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
+9 -21
View File
@@ -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