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