From 2d33d9905c47af97df7d123cd0c054b3abff3539 Mon Sep 17 00:00:00 2001 From: CaddiesNew <50736785+Nexisato@users.noreply.github.com> Date: Fri, 28 Aug 2026 15:00:49 +0800 Subject: [PATCH 01/11] feat: add multi-platform release skeleton (#1758) --- .gitignore | 4 + Makefile | 18 +++- core/cmd/swanlab-core/main.go | 164 ++++++++++++++++++++++++++++++++ core/hatch.py | 106 +++++++++++++++++++++ hatch_build.py | 110 +++++++++++++++++++++ pyproject.toml | 13 ++- scripts/build_release.sh | 174 ++++++++++++++++++++++++++++++++++ uv.lock | 83 ++++++++++++++++ 8 files changed, 670 insertions(+), 2 deletions(-) create mode 100644 core/cmd/swanlab-core/main.go create mode 100644 core/hatch.py create mode 100644 hatch_build.py create mode 100644 scripts/build_release.sh diff --git a/.gitignore b/.gitignore index cbe7fe0ad..1f022e36f 100644 --- a/.gitignore +++ b/.gitignore @@ -33,6 +33,7 @@ dist/ # testing & linting .pytest_cache/ .ruff_cache/ +.uv-cache/ .mypy_cache/ .coverage htmlcov/ @@ -66,3 +67,6 @@ CODEBUDDY.md .grok/ .reasonix/ .mcp.json + +# go core binary +swanlab/bin/ diff --git a/Makefile b/Makefile index d2df5a11a..b68b5b853 100644 --- a/Makefile +++ b/Makefile @@ -1,7 +1,7 @@ AGENTS_SKILL_DIR := .agents/skills SKILLS_DIR := docs/skills -.PHONY: init sync format proto unit bench clean build publish backport link-skills unlink-skills relink-skills core-lint core-fmt core-test core-build core-tidy +.PHONY: init sync format proto unit bench clean build publish backport link-skills unlink-skills relink-skills core-lint core-fmt core-test core-build core-tidy core-bin release-build release-verify # ---------------------------------- # SKILL (docs/skills) @@ -110,3 +110,19 @@ core-build: core-tidy: cd core && go mod tidy + +# ---------------------------------- +# Release (scripts/build_release.sh) +# ---------------------------------- + +# 本机平台编译 swanlab-core 到 swanlab/bin/,供日常联调(复用发布构建参数) +core-bin: + python3 core/hatch.py + +# 完整发布构建:sdist + any 兜底 wheel + 6 平台 wheel + 校验,全部进 dist/ +release-build: + bash scripts/build_release.sh $(VERSION) + +# 仅校验 dist/ 中已有产物(twine check + 结构完整性) +release-verify: + bash scripts/build_release.sh --verify-only diff --git a/core/cmd/swanlab-core/main.go b/core/cmd/swanlab-core/main.go new file mode 100644 index 000000000..e5bc58adc --- /dev/null +++ b/core/cmd/swanlab-core/main.go @@ -0,0 +1,164 @@ +// Command swanlab-core 是 SwanLab Go core 的进程入口。 +// +// 当前为发布脚手架阶段:仅建立进程生命周期骨架——版本上报(--version)、 +// 端点监听、父进程退出监控与信号处理;gRPC 服务端(CoreService / +// CoreSyncService / ProbeService)在后续迭代中接入,届时替换 serve 循环。 +// +// 端点约定: +// +// --listen unix:///path/to/uds Linux/macOS 进程内通信(默认路径由 Python SDK 分配) +// --listen tcp://127.0.0.1:port Windows 回环地址(uds 不可用,named pipe 支持后续提供) +// +// 生命周期:父进程退出(process 包监控)或收到 SIGINT/SIGTERM 时优雅退出, +// 防止 Python SDK 崩溃后 core 沦为孤儿进程。 +package main + +import ( + "context" + "errors" + "flag" + "fmt" + "net" + "os" + "os/signal" + "runtime" + "strconv" + "strings" + "syscall" + + "github.com/swanhubx/swanlab/core/internal/pkg/console" + "github.com/swanhubx/swanlab/core/internal/pkg/process" +) + +// version 与 commit 由构建管线通过 -ldflags -X 注入(见 core/hatch.py), +// 缺省值仅供本地 go run / go build 使用。 +var ( + version = "dev" + commit = "unknown" +) + +// 与命令行参数等价的环境变量,供 Python SDK spawn 时注入。 +const ( + envListenAddr = "SWANLAB_CORE_LISTEN" + envParentPID = "SWANLAB_CORE_PARENT_PID" +) + +// 进程退出码约定:0 正常退出(含信号触发的优雅关闭);2 用法错误;1 运行错误。 +const ( + exitUsageError = 2 + exitRunError = 1 +) + +func main() { + os.Exit(run(os.Args[1:])) +} + +func run(args []string) int { + fs := flag.NewFlagSet("swanlab-core", flag.ContinueOnError) + printVersion := fs.Bool("version", false, "打印版本信息后退出") + listenAddr := fs.String("listen", os.Getenv(envListenAddr), + "监听端点,格式 unix:// 或 tcp://<地址:端口>;Windows 仅支持 tcp:// 回环地址") + parentPID := fs.Int("parent-pid", envInt(envParentPID), + "预期父进程 PID,父进程退出时 core 随之退出;未指定时取启动瞬间的实际父进程") + if err := fs.Parse(args); err != nil { + if errors.Is(err, flag.ErrHelp) { + return 0 + } + return exitUsageError + } + + if *printVersion { + fmt.Printf("swanlab-core %s (commit %s)\n", version, commit) + return 0 + } + if *listenAddr == "" { + console.Error("未指定监听端点:通过 --listen 或环境变量 " + envListenAddr + " 传入") + return exitUsageError + } + + ln, err := listen(*listenAddr) + if err != nil { + console.Error("监听失败:", err) + return exitRunError + } + defer ln.Close() + + // 父进程监控:显式传入的 PID 优先(Python SDK 启动约定),未传时回退为 + // 监控启动瞬间的实际父进程(本地终端运行场景)。监控建立失败按约定终止启动。 + pid := *parentPID + if pid <= 0 { + pid = os.Getppid() + } + parentExited, err := process.NotifyOnParentExit(pid) + if err != nil { + console.Error("父进程监控建立失败,终止启动:", err) + return exitRunError + } + + ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM) + defer stop() + + console.Infof("swanlab-core %s listening on %s (parent pid %d)", version, *listenAddr, pid) + + serveErr := make(chan error, 1) + go func() { + serveErr <- serve(ln) + }() + + select { + case <-ctx.Done(): + console.Info("收到退出信号,正在关闭") + case <-parentExited: + console.Warning("父进程已退出,core 随之退出") + case err := <-serveErr: + if err != nil { + console.Error("监听异常退出:", err) + return exitRunError + } + } + return 0 +} + +// listen 按协议前缀创建监听器。uds 仅在非 Windows 平台可用;Windows 使用 +// TCP 回环地址兜底(named pipe 接入后在此分支扩展)。 +func listen(addr string) (net.Listener, error) { + scheme, rest, ok := strings.Cut(addr, "://") + if !ok { + return nil, fmt.Errorf("监听端点缺少协议前缀(unix:// 或 tcp://): %s", addr) + } + switch scheme { + case "unix": + if runtime.GOOS == "windows" { + return nil, errors.New("windows 平台不支持 unix:// 端点,请使用 tcp://127.0.0.1:<端口>") + } + return net.Listen("unix", rest) + case "tcp": + return net.Listen("tcp", rest) + default: + return nil, fmt.Errorf("不支持的监听协议 %q(仅 unix:// 或 tcp://)", scheme) + } +} + +// serve 接受连接后立即关闭。脚手架阶段仅验证端点连通性, +// gRPC 服务端就绪后由此接入 serve 逻辑。 +func serve(ln net.Listener) error { + for { + conn, err := ln.Accept() + if err != nil { + if errors.Is(err, net.ErrClosed) { + return nil + } + return err + } + _ = conn.Close() + } +} + +// envInt 解析整型环境变量,缺失或非法时返回 0。 +func envInt(name string) int { + v, err := strconv.Atoi(strings.TrimSpace(os.Getenv(name))) + if err != nil { + return 0 + } + return v +} diff --git a/core/hatch.py b/core/hatch.py new file mode 100644 index 000000000..8194d78bc --- /dev/null +++ b/core/hatch.py @@ -0,0 +1,106 @@ +"""swanlab-core Go 模块编译封装。 + +供两处调用: + +- ``hatch_build.py`` 构建钩子(发布/平台 wheel 构建) +- 直接执行 ``python3 core/hatch.py``(``make core-bin``,本机构建) + +职责单一:以 ``core/`` 为工作目录执行 ``go build``,将产物输出到仓库根的 +``swanlab/bin/swanlab-core(.exe)``,并注入版本号与 commit。 + +注:``CGO_ENABLED=0`` 静态编译下,Go 内部链接器写入的 ELF ``.gnu.version`` / +``.gnu.version_r`` section 会导致 auditwheel 崩溃。仅当未来引入 manylinux +容器认证(cibuildwheel 路线)时,才需要在构建流程中用 objcopy 移除这两个 +section;当前自声明 manylinux tag 的构建不涉及 auditwheel,无需处理。 +""" + +from __future__ import annotations + +import json +import os +import pathlib +import subprocess + +_REPO_ROOT = pathlib.Path(__file__).resolve().parent.parent +_CORE_DIR = _REPO_ROOT / "core" +_ENTRY_PACKAGE = "./cmd/swanlab-core" + + +def build_core( + go_binary: pathlib.Path, + output_path: pathlib.PurePath, + target_system: str | None = None, + target_arch: str | None = None, +) -> None: + """编译 swanlab-core。 + + Args: + go_binary: go 可执行文件路径,必须存在。 + output_path: 产物路径,相对仓库根(如 ``swanlab/bin/swanlab-core``)。 + target_system: 目标 GOOS,``None`` 表示使用本机平台。 + target_arch: 目标 GOARCH,``None`` 表示使用本机平台。 + """ + version = _package_version() + commit = _git_commit() + # -s -w 移除符号表与 DWARF 调试信息;-X 将版本信息注入 main 包变量 + ldflags = f"-s -w -X main.version={version} -X main.commit={commit}" + # go build 以 core/ 为工作目录,因此输出路径需要相对 core/ 前移一级 + output = pathlib.Path("..") / output_path + + subprocess.check_call( + [ + str(go_binary), + "build", + "-trimpath", + f"-ldflags={ldflags}", + "-o", + str(output), + _ENTRY_PACKAGE, + ], + cwd=str(_CORE_DIR), + env=_go_env(target_system, target_arch), + ) + + if not _is_windows_target(target_system): + # 显式设置可执行位:wheel 以 zip 外部属性记录权限,pip 解压时还原 + os.chmod(_REPO_ROOT / output_path, 0o755) + + +def _go_env(target_system: str | None, target_arch: str | None) -> dict[str, str]: + env = os.environ.copy() + if target_system: + env["GOOS"] = target_system + if target_arch: + env["GOARCH"] = target_arch + # 纯 Go 无 cgo:静态编译,保证全平台交叉编译无宿主依赖 + env["CGO_ENABLED"] = "0" + return env + + +def _is_windows_target(target_system: str | None) -> bool: + if target_system: + return target_system == "windows" + return os.name == "nt" + + +def _package_version() -> str: + """从 swanlab/package.json 读取版本号,与 SDK 版本保持同源。""" + package_json = _REPO_ROOT / "swanlab" / "package.json" + version = json.loads(package_json.read_text(encoding="utf-8"))["version"] + return str(version) + + +def _git_commit() -> str: + """返回当前 commit SHA;不在 git 仓库中或 git 不可用时返回 ``unknown``。""" + try: + return subprocess.check_output(["git", "rev-parse", "HEAD"], cwd=str(_REPO_ROOT), text=True).strip() + except Exception: + return "unknown" + + +if __name__ == "__main__": + # 本机构建入口(make core-bin),复用与发布构建一致的编译参数 + exe = "swanlab-core.exe" if os.name == "nt" else "swanlab-core" + go = pathlib.Path(os.environ.get("GO", "go")) + build_core(go_binary=go, output_path=pathlib.PurePath("swanlab", "bin", exe)) + print(f"swanlab/bin/{exe} built for host platform") diff --git a/hatch_build.py b/hatch_build.py new file mode 100644 index 000000000..b397cb951 --- /dev/null +++ b/hatch_build.py @@ -0,0 +1,110 @@ +"""SwanLab wheel 自定义构建钩子(hatchling 插件)。 + +由 ``SWANLAB_BUILD_PLATFORM`` 环境变量控制构建行为: + +- 未设置:不执行任何操作,产出纯 ``py3-none-any`` wheel。 +- 设置为平台 tag(如 ``manylinux_2_17_aarch64.manylinux2014_aarch64``、 + ``win_arm64``、``macosx_12_0_arm64``):交叉编译 Go core 内嵌进 wheel, + 并将该平台 tag 写入 wheel 元数据。 +""" + +import os +import pathlib +import re +import shutil +import sys + +from hatchling.builders.hooks.plugin.interface import BuildHookInterface + +_BUILD_PLATFORM_ENV = "SWANLAB_BUILD_PLATFORM" + +# 支持的平台 tag 形态示例(GOOS 由前缀判定,GOARCH 由结尾后缀判定) +_SUPPORTED_TAG_EXAMPLES = ( + "manylinux_2_17_x86_64.manylinux2014_x86_64", + "manylinux_2_17_aarch64.manylinux2014_aarch64", + "win_amd64", + "win_arm64", + "macosx_12_0_x86_64", + "macosx_12_0_arm64", +) + + +class CustomBuildHook(BuildHookInterface): + """按目标平台编译 Go core,并产出对应平台 tag 的 wheel。""" + + def initialize(self, version, build_data): + if self.target_name == "wheel": + self._prepare_wheel(build_data) + + def _prepare_wheel(self, build_data): + platform_tag = os.getenv(_BUILD_PLATFORM_ENV) + if not platform_tag: + return + + goos, goarch = self._parse_target_platform(platform_tag) + output = self._build_core(goos, goarch) + + build_data["tag"] = f"py3-none-{platform_tag}" + # hatchling 统一使用正斜杠路径,Windows 构建环境下亦然 + build_data["artifacts"].append(output.as_posix()) + + def _parse_target_platform(self, platform_tag): + """从平台 tag 解析 (GOOS, GOARCH),无法解析时中止构建。""" + tag = platform_tag.strip().lower() + if tag.startswith(("manylinux_", "musllinux_", "linux_")): + goos = "linux" + elif tag.startswith(("win_", "win-")): + goos = "windows" + elif tag.startswith(("macosx_", "macosx-")): + goos = "darwin" + else: + self.app.abort( + f"Unrecognized {_BUILD_PLATFORM_ENV}={platform_tag!r}: cannot determine OS prefix. " + f"Supported tags: {', '.join(_SUPPORTED_TAG_EXAMPLES)}", + ) + raise AssertionError("unreachable") + + arch_match = re.search(r"(?:-|_)(aarch64|arm64|x86_64|amd64)$", tag) + if not arch_match: + self.app.abort( + f"Cannot parse target architecture from {_BUILD_PLATFORM_ENV}={platform_tag!r} " + "(expected suffix _amd64 / _x86_64 / _arm64 / _aarch64)", + ) + raise AssertionError("unreachable") + + goarch = {"amd64": "amd64", "x86_64": "amd64", "arm64": "arm64", "aarch64": "arm64"}[arch_match.group(1)] + return goos, goarch + + def _build_core(self, goos, goarch): + """编译 Go core 到 swanlab/bin/,返回产物路径(相对仓库根)。""" + go = shutil.which("go") + if not go: + self.app.abort( + f"{_BUILD_PLATFORM_ENV} is set but no Go toolchain found. " + "Install Go to build platform wheels: https://go.dev/doc/install", + ) + raise AssertionError("unreachable") + + output = pathlib.Path("swanlab", "bin", "swanlab-core") + if goos == "windows": + output = output.with_suffix(".exe") + + self.app.display_waiting(f"Building swanlab-core ({goos}/{goarch})...") + try: + # 惰性导入:sdist 不携带 core/ 源码,仅平台构建时才需要该模块 + sys.path.insert(0, str(pathlib.Path(__file__).parent)) + from core import hatch as hatch_core + except ImportError: + self.app.abort( + f"{_BUILD_PLATFORM_ENV} is set but core/hatch.py is missing. " + "Platform wheels must be built from a git checkout; sdist does not include Go sources.", + ) + raise AssertionError("unreachable") + + hatch_core.build_core( + go_binary=pathlib.Path(go), + output_path=output, + target_system=goos, + target_arch=goarch, + ) + return output diff --git a/pyproject.toml b/pyproject.toml index c1785f252..f84aaacf3 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -80,6 +80,7 @@ s3 = [ [dependency-groups] dev = [ "build", + "hatchling>=1.18", "pytest", "pytest-mock", "pytest-benchmark", @@ -112,9 +113,19 @@ swanlab = "swanlab.cli:cli" "Documentation" = "https://docs.swanlab.cn/zh/guide_cloud/general/what-is-swanlab.html" [build-system] -requires = ["hatchling", "hatch-fancy-pypi-readme>=22.5.0"] +# hatchling>=1.18 支持自定义构建 hooks 的最低兼容版本(含 tag/artifacts ) +requires = ["hatchling>=1.18", "hatch-fancy-pypi-readme>=22.5.0"] build-backend = "hatchling.build" +# 设置 SWANLAB_BUILD_PLATFORM 时交叉编译 +[tool.hatch.build.hooks.custom] +path = "hatch_build.py" + +# sdist 不携带 Go core:从 sdist 构建恒为纯 py3-none-any wheel(core_python); +# 含二进制的平台 wheel 需从 git checkout 进行构建。 +[tool.hatch.build.targets.sdist] +exclude = ["/core"] + [tool.hatch.version] path = "swanlab/package.json" diff --git a/scripts/build_release.sh b/scripts/build_release.sh new file mode 100644 index 000000000..9f59bf0db --- /dev/null +++ b/scripts/build_release.sh @@ -0,0 +1,174 @@ +#!/usr/bin/env bash +# SwanLab 多平台 release 构建脚本。 +# +# 产物:1 sdist + 1 any 兜底 wheel + 6 平台 wheel(内嵌交叉编译的 Go core)。 +# 交叉编译与平台 tag 写入由 hatch_build.py 构建钩子完成, +# 本脚本负责切环境变量逐平台构建,以及产物完整性校验。 +# +# 用法: +# bash scripts/build_release.sh [VERSION] # 更新版本号并构建全部产物,随后校验 +# bash scripts/build_release.sh # 使用 swanlab/package.json 现有版本 +# bash scripts/build_release.sh --verify-only # 仅校验已有 dist/,不构建 +set -euo pipefail +cd "$(dirname "$0")/.." + +PLATFORMS=( + "manylinux_2_17_x86_64.manylinux2014_x86_64" # linux/amd64 + "manylinux_2_17_aarch64.manylinux2014_aarch64" # linux/arm64 + "win_amd64" # windows/amd64 + "win_arm64" # windows/arm64 + "macosx_12_0_x86_64" # darwin/amd64 + "macosx_12_0_arm64" # darwin/arm64 +) + +VERIFY_ONLY=false +VERSION="" +for arg in "$@"; do + case "$arg" in + --verify-only) VERIFY_ONLY=true ;; + -h|--help) grep '^#' "$0" | sed 's/^# \{0,1\}//'; exit 0 ;; + *) VERSION="$arg" ;; + esac +done + +if [ "$VERIFY_ONLY" = false ]; then + # 0. 版本号(沿用 make build 的写回逻辑) + if [ -n "$VERSION" ]; then + python3 -c "import json; data=json.load(open('swanlab/package.json')); data['version']='$VERSION'; json.dump(data,open('swanlab/package.json','w'),indent=2)" + echo "==> Updated swanlab/package.json version to $VERSION" + fi + + # 1. 清理旧产物 + rm -rf dist swanlab/bin + + # 2. sdist + any 兜底 wheel(不设 SWANLAB_BUILD_PLATFORM,钩子零动作) + echo "==> Building sdist + fallback any wheel" + uv build --sdist + uv build --wheel + + # 3. 逐平台交叉编译(钩子内完成:go build + 原生平台 tag + artifacts 注入) + for TAG in "${PLATFORMS[@]}"; do + echo "==> Building platform wheel: $TAG" + SWANLAB_BUILD_PLATFORM="$TAG" uv build --wheel + rm -rf swanlab/bin # 清理,避免污染下一平台 + done +fi + +# 4. 校验 +echo "==> twine check" +uvx twine check dist/* + +echo "==> 结构校验" +python3 - <<'PY' +import pathlib +import re +import subprocess +import sys +import tarfile +import tempfile +import zipfile + +DIST = pathlib.Path("dist") + +# tag -> (二进制文件名, file 输出需命中的关键词,任一命中即可) +PLATFORMS = { + "manylinux_2_17_x86_64.manylinux2014_x86_64": ("swanlab-core", ["elf", "x86-64"]), + "manylinux_2_17_aarch64.manylinux2014_aarch64": ("swanlab-core", ["elf", "aarch64"]), + "win_amd64": ("swanlab-core.exe", ["pe32+", "x86-64", "x86_64"]), + "win_arm64": ("swanlab-core.exe", ["pe32+", "aarch64", "arm64"]), + "macosx_12_0_x86_64": ("swanlab-core", ["mach-o", "x86_64"]), + "macosx_12_0_arm64": ("swanlab-core", ["mach-o", "arm64"]), +} + +failures = [] + + +def check(cond, msg): + print(("PASS " if cond else "FAIL ") + msg) + if not cond: + failures.append(msg) + + +def wheel_version(name): + m = re.fullmatch(r"swanlab-(.+)-py3-none-(.+)\.whl", name) + return (m.group(1), m.group(2)) if m else None + + +wheels = sorted(DIST.glob("swanlab-*-py3-none-*.whl")) +sdists = sorted(DIST.glob("swanlab-*.tar.gz")) +parsed = [wheel_version(w.name) for w in wheels] +check(all(p is not None for p in parsed), "wheel 文件名均可解析") +check(len(sdists) == 1, f"恰好 1 个 sdist(实际 {len(sdists)})") + +# 版本一致性:所有产物(含 sdist)使用同一版本串 +versions = {p[0] for p in parsed if p} | {re.fullmatch(r"swanlab-(.+)\.tar\.gz", s.name).group(1) for s in sdists} +check(len(versions) == 1, f"全部产物版本一致(实际 {versions})") + +# 产物齐全:any 兜底 + 6 平台 wheel,无多余文件 +tags = {p[1] for p in parsed if p} +check(tags == set(PLATFORMS) | {"any"}, f"wheel tag 齐全(多出: {sorted(tags - set(PLATFORMS) - {'any'})}, 缺少: {sorted(set(PLATFORMS) - tags)})") + +for whl in wheels: + tag = wheel_version(whl.name)[1] + with zipfile.ZipFile(whl) as zf: + names = zf.namelist() + has_binary = any(n.startswith("swanlab/bin/") for n in names) + + if tag == "any": + check(not has_binary, "any wheel 不含二进制") + continue + + exe = PLATFORMS[tag][0] + entry = f"swanlab/bin/{exe}" + if entry not in names: + check(False, f"{tag}: 缺少 {entry}") + continue + + mode = zf.getinfo(entry).external_attr >> 16 + check(mode & 0o111 != 0, f"{tag}: 二进制含可执行位") + + with tempfile.TemporaryDirectory() as td: + extracted = zf.extract(entry, td) + desc = subprocess.run(["file", "-b", str(extracted)], capture_output=True, text=True).stdout.lower() + check(any(k in desc for k in PLATFORMS[tag][1]), f"{tag}: file 架构匹配({desc.strip()})") + + if tag.startswith("manylinux"): + # glibc tag 双写:WHEEL 元数据需逐条展开为多行 Tag: + meta_name = next(n for n in names if n.endswith(".dist-info/WHEEL")) + meta = zf.read(meta_name).decode() + lines = [ln.split(": ", 1)[1] for ln in meta.splitlines() if ln.startswith("Tag:")] + expected = {f"py3-none-{t}" for t in tag.split(".")} + check(expected <= set(lines), f"{tag}: WHEEL 元数据多行 Tag 展开(实际 {lines})") + +# sdist:不含 core/ 源码与二进制,但必须携带构建钩子(保证从 sdist 重建 wheel 可行) +with tarfile.open(sdists[0]) as tf: + members = tf.getnames() +check(any(m.count("/") == 1 and m.endswith("/hatch_build.py") for m in members), "sdist 携带 hatch_build.py") +check(not any(m.count("/") == 1 and m.split("/")[1] == "core" for m in members), "sdist 不含 core/ 源码") +check(not any("/swanlab/bin/" in m for m in members), "sdist 不含二进制") + + +def py_files(wheel): + with zipfile.ZipFile(wheel) as zf: + return {n for n in zf.namelist() if n.endswith(".py")} + + +# sdist 可重建性:从 sdist 重建 wheel,.py 清单须与 checkout 直建的 any wheel 一致 +# (防止 sdist exclude 误伤嵌套同名目录) +any_wheel = next(w for w in wheels if w.name.endswith("-py3-none-any.whl")) +with tempfile.TemporaryDirectory() as td: + with tarfile.open(sdists[0]) as tf: + tf.extractall(td) + src = next(pathlib.Path(td).glob("swanlab-*/")) + subprocess.run(["uv", "build", "--wheel", "--out-dir", td], cwd=src, check=True, capture_output=True) + rebuilt = next(pathlib.Path(td).glob("*.whl")) + diff = sorted(py_files(any_wheel) ^ py_files(rebuilt)) + check(not diff, f"sdist 重建 wheel 与 any wheel 的 .py 清单一致(差异 {len(diff)} 项: {diff[:5]})") + +if failures: + print(f"\n{len(failures)} 项校验失败", file=sys.stderr) + sys.exit(1) +print("\nAll checks passed.") +PY + +echo "==> Done" diff --git a/uv.lock b/uv.lock index 7af16adad..18aade307 100644 --- a/uv.lock +++ b/uv.lock @@ -1273,6 +1273,59 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/04/4b/29cac41a4d98d144bf5f6d33995617b185d14b22401f75ca86f384e87ff1/h11-0.16.0-py3-none-any.whl", hash = "sha256:63cf8bbe7522de3bf65932fda1d9c2772064ffb3dae62d55932da54b31cb6c86", size = 37515, upload-time = "2025-04-24T03:35:24.344Z" }, ] +[[package]] +name = "hatchling" +version = "1.27.0" +source = { registry = "https://pypi.org/simple" } +resolution-markers = [ + "python_full_version < '3.10' and sys_platform == 'linux'", + "python_full_version < '3.10' and sys_platform != 'linux'", +] +dependencies = [ + { name = "packaging" }, + { name = "pathspec" }, + { name = "pluggy" }, + { name = "tomli" }, + { name = "trove-classifiers" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/8f/8a/cc1debe3514da292094f1c3a700e4ca25442489731ef7c0814358816bb03/hatchling-1.27.0.tar.gz", hash = "sha256:971c296d9819abb3811112fc52c7a9751c8d381898f36533bb16f9791e941fd6", size = 54983, upload-time = "2024-12-15T17:08:11.894Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/08/e7/ae38d7a6dfba0533684e0b2136817d667588ae3ec984c1a4e5df5eb88482/hatchling-1.27.0-py3-none-any.whl", hash = "sha256:d3a2f3567c4f926ea39849cdf924c7e99e6686c9c8e288ae1037c8fa2a5d937b", size = 75794, upload-time = "2024-12-15T17:08:10.364Z" }, +] + +[[package]] +name = "hatchling" +version = "1.32.0" +source = { registry = "https://pypi.org/simple" } +resolution-markers = [ + "python_full_version >= '3.14' and sys_platform == 'linux'", + "python_full_version >= '3.14' and sys_platform == 'win32'", + "python_full_version >= '3.14' and sys_platform == 'emscripten'", + "python_full_version >= '3.14' and sys_platform != 'emscripten' and sys_platform != 'linux' and sys_platform != 'win32'", + "python_full_version >= '3.12' and python_full_version < '3.14' and sys_platform == 'linux'", + "python_full_version >= '3.12' and python_full_version < '3.14' and sys_platform == 'win32'", + "python_full_version >= '3.12' and python_full_version < '3.14' and sys_platform == 'emscripten'", + "python_full_version >= '3.12' and python_full_version < '3.14' and sys_platform != 'emscripten' and sys_platform != 'linux' and sys_platform != 'win32'", + "python_full_version == '3.11.*' and sys_platform == 'linux'", + "python_full_version == '3.11.*' and sys_platform == 'win32'", + "python_full_version == '3.11.*' and sys_platform == 'emscripten'", + "python_full_version == '3.11.*' and sys_platform != 'emscripten' and sys_platform != 'linux' and sys_platform != 'win32'", + "python_full_version == '3.10.*' and sys_platform == 'linux'", + "python_full_version == '3.10.*' and sys_platform != 'linux'", +] +dependencies = [ + { name = "packaging" }, + { name = "pathspec" }, + { name = "pluggy" }, + { name = "tomli", marker = "python_full_version < '3.11'" }, + { name = "tomlkit" }, + { name = "trove-classifiers" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/69/08/33331757185504aae48b8d9bd78cec03a76e3aecfb52e549d05a2347c0dd/hatchling-1.32.0.tar.gz", hash = "sha256:0bdbde4a52b06c37e3eca395f85a762bf0ef06fe374fd8ae429dc6be10230f5f", size = 57783, upload-time = "2026-08-11T05:03:44.114Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/a9/84/1798b6d85ecde0e31546004efd25c5de1b1f49250644a60cce460e12593a/hatchling-1.32.0-py3-none-any.whl", hash = "sha256:0e17c9c3b9aa7c625acc8d0f5b622f107d5049af9ecf5ada4de1aada5be7cdbc", size = 78435, upload-time = "2026-08-11T05:03:42.644Z" }, +] + [[package]] name = "identify" version = "2.6.15" @@ -3189,6 +3242,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/b6/61/fae042894f4296ec49e3f193aff5d7c18440da9e48102c3315e1bc4519a7/parso-0.8.6-py2.py3-none-any.whl", hash = "sha256:2c549f800b70a5c4952197248825584cb00f033b29c692671d3bf08bf380baff", size = 106894, upload-time = "2026-02-09T15:45:21.391Z" }, ] +[[package]] +name = "pathspec" +version = "1.1.1" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/5a/82/42f767fc1c1143d6fd36efb827202a2d997a375e160a71eb2888a925aac1/pathspec-1.1.1.tar.gz", hash = "sha256:17db5ecd524104a120e173814c90367a96a98d07c45b2e10c2f3919fff91bf5a", size = 135180, upload-time = "2026-04-27T01:46:08.907Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/f1/d9/7fb5aa316bc299258e68c73ba3bddbc499654a07f151cba08f6153988714/pathspec-1.1.1-py3-none-any.whl", hash = "sha256:a00ce642f577bf7f473932318056212bc4f8bfdf53128c78bbd5af0b9b20b189", size = 57328, upload-time = "2026-04-27T01:46:07.06Z" }, +] + [[package]] name = "peewee" version = "3.19.0" @@ -5227,6 +5289,8 @@ dev = [ { name = "build" }, { name = "freezegun" }, { name = "grpcio-tools" }, + { name = "hatchling", version = "1.27.0", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.10'" }, + { name = "hatchling", version = "1.32.0", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.10'" }, { name = "ipykernel", version = "6.31.0", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.10'" }, { name = "ipykernel", version = "7.2.0", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.10'" }, { name = "ipython", version = "8.18.1", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.10'" }, @@ -5300,6 +5364,7 @@ dev = [ { name = "build" }, { name = "freezegun" }, { name = "grpcio-tools" }, + { name = "hatchling", specifier = ">=1.18" }, { name = "ipykernel" }, { name = "ipython" }, { name = "mmengine" }, @@ -5432,6 +5497,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/23/d1/136eb2cb77520a31e1f64cbae9d33ec6df0d78bdf4160398e86eec8a8754/tomli-2.4.0-py3-none-any.whl", hash = "sha256:1f776e7d669ebceb01dee46484485f43a4048746235e683bcdffacdf1fb4785a", size = 14477, upload-time = "2026-01-11T11:22:37.446Z" }, ] +[[package]] +name = "tomlkit" +version = "0.15.1" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/94/96/e07752635b98536177fa1f37671c8f3cdde2e724c6bcf6034b2cfb571565/tomlkit-0.15.1.tar.gz", hash = "sha256:e25bbf38843005246210a12982776f27f99cb9be67160e14434d0c0d21ee1e97", size = 180129, upload-time = "2026-07-17T01:48:04.562Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/13/bc/8c13eb66537dce1d2bd3a57132902f38d0e7f5bb46fa9f4daed9fe9d76ee/tomlkit-0.15.1-py3-none-any.whl", hash = "sha256:177a05aece5a8ca5266fd3c448abb47b8d352f09d477d3ca8332db4d89b24304", size = 49449, upload-time = "2026-07-17T01:48:05.728Z" }, +] + [[package]] name = "torch" version = "2.8.0" @@ -5756,6 +5830,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/f6/56/6113c23ff46c00aae423333eb58b3e60bdfe9179d542781955a5e1514cb3/triton-3.6.0-cp314-cp314t-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:46bd1c1af4b6704e554cad2eeb3b0a6513a980d470ccfa63189737340c7746a7", size = 188397994, upload-time = "2026-01-20T16:01:14.236Z" }, ] +[[package]] +name = "trove-classifiers" +version = "2026.6.1.19" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/c2/e3/7ca82ee24c82d344584abd5b8637b3bd056f2900226e8d82fc22f1184b92/trove_classifiers-2026.6.1.19.tar.gz", hash = "sha256:c5132b4b61a829d11cfbd2d72e97f20a45ed6edb95e45c5efdeb5e00836b2745", size = 17059, upload-time = "2026-06-01T19:41:34.649Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/7c/a4/81502f486f01db95bc8320646a8a12511f5e556cb63d5e224d91816605c4/trove_classifiers-2026.6.1.19-py3-none-any.whl", hash = "sha256:ab4c4ec93cc4a4e7815fa759906e05e6bb3f2fbd92ea0f897288c6a43efd15b3", size = 14211, upload-time = "2026-06-01T19:41:33.434Z" }, +] + [[package]] name = "typing-extensions" version = "4.15.0" From 27a669b1495b099df011fb59e740643776e72040 Mon Sep 17 00:00:00 2001 From: Kang Li <79990647+SAKURA-CAT@users.noreply.github.com> Date: Thu, 27 Aug 2026 13:58:03 +0800 Subject: [PATCH 02/11] chore: add docs/plans directory to gitignore (#1759) --- .gitignore | 1 + 1 file changed, 1 insertion(+) diff --git a/.gitignore b/.gitignore index 1f022e36f..4f2bbf9f5 100644 --- a/.gitignore +++ b/.gitignore @@ -61,6 +61,7 @@ GEMINI.md QWEN.md .omx/ .opencode/ +docs/plans CODEBUDDY.md .codegraph/ .mimocode/ From 5ed3229a8970005c045540dae2de493e38701f2e Mon Sep 17 00:00:00 2001 From: Kang Li <79990647+SAKURA-CAT@users.noreply.github.com> Date: Fri, 28 Aug 2026 15:06:20 +0800 Subject: [PATCH 03/11] chore: update author emails in project metadata (#1761) --- pyproject.toml | 9 +++++---- 1 file changed, 5 insertions(+), 4 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index f84aaacf3..3655efb93 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -5,10 +5,11 @@ license = "Apache-2.0" description = "Python library for streamlined tracking and management of AI training processes." requires-python = ">=3.9" authors = [ - { name = "Cunyue", email = "team@swanhub.co" }, - { name = "ZeYi Lin", email = "team@swanhub.co" }, - { name = "Kaikaikaifang", email = "team@swanhub.co" }, - { name = "Feudalman", email = "team@swanhub.co" }, + { name = "Cunyue", email = "kang.li@emotionmachine.cn" }, + { name = "CaddiesNew", email = "shaobo.niu@emotionmachine.cn" }, + { name = "ZeYi Lin", email = "zeyi.lin@emotionmachine.cn" }, + { name = "Kaikaikaifang", email = "kaifang.ji@emotionmachine.cn" }, + { name = "Feudalman", email = "zirui.cai@emotionmachine.cn" }, ] keywords = [ "machine learning", From 85092a00d43f4c2f5893b473c12beb0b6fb5af36 Mon Sep 17 00:00:00 2001 From: CaddiesNew <50736785+Nexisato@users.noreply.github.com> Date: Sat, 29 Aug 2026 20:10:43 +0800 Subject: [PATCH 04/11] feat: add define_metric API for chart customization (#1733) Co-authored-by: Kang Li --- Makefile | 1 + .../swanlab/metric/column/v1/column.pb.go | 31 +- protos/swanlab/metric/column/v1/column.proto | 8 + swanlab/__init__.py | 4 +- swanlab/__init__.pyi | 68 ++- .../swanlab/metric/column/v1/column_pb2.py | 20 +- .../swanlab/metric/column/v1/column_pb2.pyi | 8 +- swanlab/sdk/__init__.py | 4 +- swanlab/sdk/cmd/run.py | 2 +- swanlab/sdk/internal/bus/__init__.py | 6 +- swanlab/sdk/internal/bus/events.py | 41 +- .../internal/core_python/transport/sender.py | 12 +- .../sdk/internal/pkg/constraints/__init__.py | 4 +- swanlab/sdk/internal/pkg/helper/__init__.py | 3 +- .../pkg/helper/{system.py => metric.py} | 18 +- .../internal/probe_python/typings/__init__.py | 4 +- swanlab/sdk/internal/run/__init__.py | 192 ++++++-- .../sdk/internal/run/components/__init__.py | 14 +- .../run/components/consumer/__init__.py | 194 +++++++- .../{ => consumer}/builder/__init__.py | 39 +- .../components/consumer/resolver/__init__.py | 342 ++++++++++++++ .../run/components/consumer/resolver/state.py | 48 ++ swanlab/sdk/internal/run/fmt.py | 19 +- swanlab/sdk/typings/core_python/api/upload.py | 2 + .../internal/core_python/test_core_sync.py | 60 ++- .../core_python/transport/test_sender.py | 148 +++++- .../{ => consumer/builder}/test_builder.py | 34 +- .../consumer/resolver/test_resolver.py | 281 +++++++++++ .../components/consumer/test_consumer_log.py | 95 ++++ .../{ => consumer}/test_consumer_save.py | 5 +- .../consumer/test_consumer_step_sync.py | 436 ++++++++++++++++++ .../internal/run/test_run_define_metric.py | 100 ++++ 32 files changed, 2048 insertions(+), 195 deletions(-) rename swanlab/sdk/internal/pkg/helper/{system.py => metric.py} (50%) rename swanlab/sdk/internal/run/components/{ => consumer}/builder/__init__.py (85%) create mode 100644 swanlab/sdk/internal/run/components/consumer/resolver/__init__.py create mode 100644 swanlab/sdk/internal/run/components/consumer/resolver/state.py rename tests/unit/sdk/internal/run/components/{ => consumer/builder}/test_builder.py (74%) create mode 100644 tests/unit/sdk/internal/run/components/consumer/resolver/test_resolver.py create mode 100644 tests/unit/sdk/internal/run/components/consumer/test_consumer_log.py rename tests/unit/sdk/internal/run/components/{ => consumer}/test_consumer_save.py (95%) create mode 100644 tests/unit/sdk/internal/run/components/consumer/test_consumer_step_sync.py create mode 100644 tests/unit/sdk/internal/run/test_run_define_metric.py diff --git a/Makefile b/Makefile index b68b5b853..94c1fa915 100644 --- a/Makefile +++ b/Makefile @@ -58,6 +58,7 @@ clean: # ---------------------------------- # Python package (swanlab/) # ---------------------------------- +.PHONY: init sync format proto unit bench clean build publish link-skills unlink-skills relink-skills init: uv sync --all-extras diff --git a/core/proto/swanlab/metric/column/v1/column.pb.go b/core/proto/swanlab/metric/column/v1/column.pb.go index 048f77c99..e2a684585 100644 --- a/core/proto/swanlab/metric/column/v1/column.pb.go +++ b/core/proto/swanlab/metric/column/v1/column.pb.go @@ -418,7 +418,13 @@ type ColumnRecord struct { // 指标名称,仅对单实验图表生效 MetricName string `protobuf:"bytes,11,opt,name=metric_name,json=metricName,proto3" json:"metric_name,omitempty"` // 指标颜色列表,十六进制颜色值,如 "#FF5733" - MetricColors *MetricColors `protobuf:"bytes,12,opt,name=metric_colors,json=metricColors,proto3" json:"metric_colors,omitempty"` + MetricColors *MetricColors `protobuf:"bytes,12,opt,name=metric_colors,json=metricColors,proto3" json:"metric_colors,omitempty"` + // define_metric: 自定义 X 轴 key;none 表示系统默认(_step) + // media 类型图表暂时无视此字段 + XAxis *string `protobuf:"bytes,13,opt,name=x_axis,json=xAxis,proto3,oneof" json:"x_axis,omitempty"` + // define_metric: 自动建图表是否放入 HIDDEN section + // SDK 上行仅在 True 时携带(presence 编码);first-log 后不可变更 + Hidden *bool `protobuf:"varint,14,opt,name=hidden,proto3,oneof" json:"hidden,omitempty"` unknownFields protoimpl.UnknownFields sizeCache protoimpl.SizeCache } @@ -537,6 +543,20 @@ func (x *ColumnRecord) GetMetricColors() *MetricColors { return nil } +func (x *ColumnRecord) GetXAxis() string { + if x != nil && x.XAxis != nil { + return *x.XAxis + } + return "" +} + +func (x *ColumnRecord) GetHidden() bool { + if x != nil && x.Hidden != nil { + return *x.Hidden + } + return false +} + var File_swanlab_metric_column_v1_column_proto protoreflect.FileDescriptor const file_swanlab_metric_column_v1_column_proto_rawDesc = "" + @@ -549,7 +569,7 @@ const file_swanlab_metric_column_v1_column_proto_rawDesc = "" + "\x04_max\"8\n" + "\fMetricColors\x12\x14\n" + "\x05light\x18\x01 \x01(\tR\x05light\x12\x12\n" + - "\x04dark\x18\x02 \x01(\tR\x04dark\"\xf9\x04\n" + + "\x04dark\x18\x02 \x01(\tR\x04dark\"\xc8\x05\n" + "\fColumnRecord\x12H\n" + "\fcolumn_class\x18\x01 \x01(\x0e2%.swanlab.metric.column.v1.ColumnClassR\vcolumnClass\x12E\n" + "\vcolumn_type\x18\x02 \x01(\x0e2$.swanlab.metric.column.v1.ColumnTypeR\n" + @@ -570,7 +590,11 @@ const file_swanlab_metric_column_v1_column_proto_rawDesc = "" + " \x01(\x0e2#.swanlab.metric.column.v1.ChartTypeR\tchartType\x12\x1f\n" + "\vmetric_name\x18\v \x01(\tR\n" + "metricName\x12K\n" + - "\rmetric_colors\x18\f \x01(\v2&.swanlab.metric.column.v1.MetricColorsR\fmetricColors*\xff\x01\n" + + "\rmetric_colors\x18\f \x01(\v2&.swanlab.metric.column.v1.MetricColorsR\fmetricColors\x12\x1a\n" + + "\x06x_axis\x18\r \x01(\tH\x00R\x05xAxis\x88\x01\x01\x12\x1b\n" + + "\x06hidden\x18\x0e \x01(\bH\x01R\x06hidden\x88\x01\x01B\t\n" + + "\a_x_axisB\t\n" + + "\a_hidden*\xff\x01\n" + "\n" + "ColumnType\x12\x1b\n" + "\x17COLUMN_TYPE_UNSPECIFIED\x10\x00\x12\x16\n" + @@ -651,6 +675,7 @@ func file_swanlab_metric_column_v1_column_proto_init() { return } file_swanlab_metric_column_v1_column_proto_msgTypes[0].OneofWrappers = []any{} + file_swanlab_metric_column_v1_column_proto_msgTypes[2].OneofWrappers = []any{} type x struct{} out := protoimpl.TypeBuilder{ File: protoimpl.DescBuilder{ diff --git a/protos/swanlab/metric/column/v1/column.proto b/protos/swanlab/metric/column/v1/column.proto index 3a4c0befb..2c73ce23c 100644 --- a/protos/swanlab/metric/column/v1/column.proto +++ b/protos/swanlab/metric/column/v1/column.proto @@ -110,4 +110,12 @@ message ColumnRecord { // 指标颜色列表,十六进制颜色值,如 "#FF5733" MetricColors metric_colors = 12; + + // define_metric: 自定义 X 轴 key;none 表示系统默认(_step) + // media 类型图表暂时无视此字段 + optional string x_axis = 13; + + // define_metric: 自动建图表是否放入 HIDDEN section + // SDK 上行仅在 True 时携带(presence 编码);first-log 后不可变更 + optional bool hidden = 14; } diff --git a/swanlab/__init__.py b/swanlab/__init__.py index 3681c6293..0a3f50e63 100644 --- a/swanlab/__init__.py +++ b/swanlab/__init__.py @@ -15,7 +15,7 @@ Video, async_log, config, - define_scalar, + define_metric, echarts, finish, get_run, @@ -64,7 +64,7 @@ "log_html", "log_object3d", "log_molecule", - "define_scalar", + "define_metric", "save", "async_log", "sync", diff --git a/swanlab/__init__.pyi b/swanlab/__init__.pyi index 5cafead5c..0e93f5a46 100644 --- a/swanlab/__init__.pyi +++ b/swanlab/__init__.pyi @@ -32,7 +32,6 @@ from .sdk import ( from .sdk.typings.cmd import ConfigLike, LoginType from .sdk.typings.context import CallbacksType from .sdk.typings.run import AsyncLogType, FinishType, ModeType, ParallelType, ResumeType, SaveType -from .sdk.typings.run.column import ScalarXAxisType from .sdk.typings.run.transforms import CaptionsType from .sdk.typings.run.transforms.audio import AudioDatasType, AudioRatesType from .sdk.typings.run.transforms.echarts import EChartsDatasType @@ -62,7 +61,7 @@ __all__ = [ "log_html", "log_object3d", "log_molecule", - "define_scalar", + "define_metric", "async_log", "save", "sync", @@ -687,33 +686,58 @@ def async_log( """ ... -def define_scalar( - *, +def define_metric( key: str, - name: Optional[str] = None, - color: Optional[str] = None, - x_axis: Optional[ScalarXAxisType] = None, - chart_name: Optional[str] = None, + *, + x_axis: Optional[str] = None, + section_name: Optional[str] = None, + hidden: Optional[bool] = None, + step_sync: Optional[bool] = None, + overwrite: bool = False, + **kwargs: Any, ) -> None: - """Explicitly define a scalar column. - - Call this before logging to customize how a scalar metric is displayed, - such as setting a display name, color, or x-axis type. - - :param key: The key for the scalar column. Supports glob patterns (e.g. "train/*") to match multiple columns at once. - :param name: Optional display name for the scalar column. - :param color: Optional hex color for the scalar line in charts. - :param x_axis: Optional x-axis type. One of "_step", "_relative_time", or a custom key. - :param chart_name: Optional name for the chart group this column belongs to. - :raises RuntimeError: If called without an active run. + """Define a metric's display configuration before logging. + + Customizes how an auto-generated chart for *key* appears in project Views: + X-axis, section placement, and visibility. The same ``(class, key)`` + shares one chart across all runs in the project. + + :param key: Metric key. Supports exact match and a single trailing ``*`` + glob (e.g. ``"train/*"``). System keys are never matched. + :param x_axis: Custom X-axis key; ``None`` means the system step. + ``step_metric`` is accepted as a ``**kwargs`` alias. + :param section_name: Section name for the auto chart. ``None`` means the + default section derived from the key. + :param hidden: If ``True``, place the chart in the HIDDEN section. + Three states: ``None`` (default) means "not provided" — merge mode + keeps the previous value; ``True`` hides; ``False`` explicitly + unhides (also effective in merge mode). + :param step_sync: When X and Y are logged separately, whether to reuse the + latest X value on the current step. Defaults to ``True`` when + ``x_axis`` (or ``step_metric``) is set, ``False`` otherwise. + :param overwrite: If ``False`` (default), merge with previous calls for + the same ``key``. If ``True``, unspecified fields reset to default. + + .. note:: + Project-wide first-writer-wins: the chart for ``(class, key)`` is + shared across all runs, and ``define_metric`` only takes effect on + the first run that introduces the key to the project. Once the chart + exists, later calls (including ``overwrite=True``) have no effect. Examples: + Custom X-axis with separate logging: + >>> import swanlab >>> swanlab.init(mode="local") - >>> swanlab.define_scalar("loss", color="#FF5733", x_axis="_step") - >>> swanlab.log({"loss": 0.5}) - >>> swanlab.finish() + >>> run = swanlab.get_run() + >>> run.define_metric("train/loss", x_axis="train/epoch") + >>> swanlab.log({"train/epoch": 1}) # step=1 + >>> swanlab.log({"train/loss": 0.8}) # step=2, auto-syncs train/epoch=1 + + Glob pattern for all validation metrics: + + >>> run.define_metric("val/*", section_name="Validation") """ ... diff --git a/swanlab/proto/swanlab/metric/column/v1/column_pb2.py b/swanlab/proto/swanlab/metric/column/v1/column_pb2.py index 75c7e31a3..d3ff3819d 100644 --- a/swanlab/proto/swanlab/metric/column/v1/column_pb2.py +++ b/swanlab/proto/swanlab/metric/column/v1/column_pb2.py @@ -24,7 +24,7 @@ -DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(b'\n%swanlab/metric/column/v1/column.proto\x12\x18swanlab.metric.column.v1\"<\n\x06YRange\x12\x10\n\x03min\x18\x01 \x01(\x01H\x00\x88\x01\x01\x12\x10\n\x03max\x18\x02 \x01(\x01H\x01\x88\x01\x01\x42\x06\n\x04_minB\x06\n\x04_max\"+\n\x0cMetricColors\x12\r\n\x05light\x18\x01 \x01(\t\x12\x0c\n\x04\x64\x61rk\x18\x02 \x01(\t\"\xeb\x03\n\x0c\x43olumnRecord\x12;\n\x0c\x63olumn_class\x18\x01 \x01(\x0e\x32%.swanlab.metric.column.v1.ColumnClass\x12\x39\n\x0b\x63olumn_type\x18\x02 \x01(\x0e\x32$.swanlab.metric.column.v1.ColumnType\x12\x12\n\ncolumn_key\x18\x03 \x01(\t\x12\x13\n\x0b\x63olumn_name\x18\x04 \x01(\t\x12\x14\n\x0csection_name\x18\x05 \x01(\t\x12;\n\x0csection_type\x18\x06 \x01(\x0e\x32%.swanlab.metric.column.v1.SectionType\x12\x31\n\x07y_range\x18\x07 \x01(\x0b\x32 .swanlab.metric.column.v1.YRange\x12\x13\n\x0b\x63hart_index\x18\x08 \x01(\t\x12\x12\n\nchart_name\x18\t \x01(\t\x12\x37\n\nchart_type\x18\n \x01(\x0e\x32#.swanlab.metric.column.v1.ChartType\x12\x13\n\x0bmetric_name\x18\x0b \x01(\t\x12=\n\rmetric_colors\x18\x0c \x01(\x0b\x32&.swanlab.metric.column.v1.MetricColors*\xff\x01\n\nColumnType\x12\x1b\n\x17\x43OLUMN_TYPE_UNSPECIFIED\x10\x00\x12\x16\n\x12\x43OLUMN_TYPE_SCALAR\x10\x01\x12\x15\n\x11\x43OLUMN_TYPE_IMAGE\x10\x02\x12\x15\n\x11\x43OLUMN_TYPE_AUDIO\x10\x03\x12\x14\n\x10\x43OLUMN_TYPE_TEXT\x10\x04\x12\x15\n\x11\x43OLUMN_TYPE_VIDEO\x10\x05\x12\x17\n\x13\x43OLUMN_TYPE_ECHARTS\x10\x06\x12\x18\n\x14\x43OLUMN_TYPE_OBJECT3D\x10\x07\x12\x18\n\x14\x43OLUMN_TYPE_MOLECULE\x10\x08\x12\x14\n\x10\x43OLUMN_TYPE_HTML\x10\t*]\n\x0b\x43olumnClass\x12\x1c\n\x18\x43OLUMN_CLASS_UNSPECIFIED\x10\x00\x12\x17\n\x13\x43OLUMN_CLASS_CUSTOM\x10\x01\x12\x17\n\x13\x43OLUMN_CLASS_SYSTEM\x10\x02*\x9d\x02\n\tChartType\x12\x1a\n\x16\x43HART_TYPE_UNSPECIFIED\x10\x00\x12\x13\n\x0f\x43HART_TYPE_LINE\x10\x01\x12\x12\n\x0e\x43HART_TYPE_BAR\x10\x02\x12\x15\n\x11\x43HART_TYPE_SCALAR\x10\x03\x12\x14\n\x10\x43HART_TYPE_IMAGE\x10\x04\x12\x14\n\x10\x43HART_TYPE_AUDIO\x10\x05\x12\x13\n\x0f\x43HART_TYPE_TEXT\x10\x06\x12\x14\n\x10\x43HART_TYPE_VIDEO\x10\x07\x12\x16\n\x12\x43HART_TYPE_ECHARTS\x10\x08\x12\x17\n\x13\x43HART_TYPE_OBJECT3D\x10\t\x12\x17\n\x13\x43HART_TYPE_MOLECULE\x10\n\x12\x13\n\x0f\x43HART_TYPE_HTML\x10\x0b*\x8f\x01\n\x0bSectionType\x12\x1c\n\x18SECTION_TYPE_UNSPECIFIED\x10\x00\x12\x17\n\x13SECTION_TYPE_PINNED\x10\x01\x12\x17\n\x13SECTION_TYPE_HIDDEN\x10\x02\x12\x17\n\x13SECTION_TYPE_PUBLIC\x10\x03\x12\x17\n\x13SECTION_TYPE_SYSTEM\x10\x04\x42JZHgithub.com/swanhubx/swanlab/core/proto/swanlab/metric/column/v1;columnv1b\x06proto3') +DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(b'\n%swanlab/metric/column/v1/column.proto\x12\x18swanlab.metric.column.v1\"<\n\x06YRange\x12\x10\n\x03min\x18\x01 \x01(\x01H\x00\x88\x01\x01\x12\x10\n\x03max\x18\x02 \x01(\x01H\x01\x88\x01\x01\x42\x06\n\x04_minB\x06\n\x04_max\"+\n\x0cMetricColors\x12\r\n\x05light\x18\x01 \x01(\t\x12\x0c\n\x04\x64\x61rk\x18\x02 \x01(\t\"\xab\x04\n\x0c\x43olumnRecord\x12;\n\x0c\x63olumn_class\x18\x01 \x01(\x0e\x32%.swanlab.metric.column.v1.ColumnClass\x12\x39\n\x0b\x63olumn_type\x18\x02 \x01(\x0e\x32$.swanlab.metric.column.v1.ColumnType\x12\x12\n\ncolumn_key\x18\x03 \x01(\t\x12\x13\n\x0b\x63olumn_name\x18\x04 \x01(\t\x12\x14\n\x0csection_name\x18\x05 \x01(\t\x12;\n\x0csection_type\x18\x06 \x01(\x0e\x32%.swanlab.metric.column.v1.SectionType\x12\x31\n\x07y_range\x18\x07 \x01(\x0b\x32 .swanlab.metric.column.v1.YRange\x12\x13\n\x0b\x63hart_index\x18\x08 \x01(\t\x12\x12\n\nchart_name\x18\t \x01(\t\x12\x37\n\nchart_type\x18\n \x01(\x0e\x32#.swanlab.metric.column.v1.ChartType\x12\x13\n\x0bmetric_name\x18\x0b \x01(\t\x12=\n\rmetric_colors\x18\x0c \x01(\x0b\x32&.swanlab.metric.column.v1.MetricColors\x12\x13\n\x06x_axis\x18\r \x01(\tH\x00\x88\x01\x01\x12\x13\n\x06hidden\x18\x0e \x01(\x08H\x01\x88\x01\x01\x42\t\n\x07_x_axisB\t\n\x07_hidden*\xff\x01\n\nColumnType\x12\x1b\n\x17\x43OLUMN_TYPE_UNSPECIFIED\x10\x00\x12\x16\n\x12\x43OLUMN_TYPE_SCALAR\x10\x01\x12\x15\n\x11\x43OLUMN_TYPE_IMAGE\x10\x02\x12\x15\n\x11\x43OLUMN_TYPE_AUDIO\x10\x03\x12\x14\n\x10\x43OLUMN_TYPE_TEXT\x10\x04\x12\x15\n\x11\x43OLUMN_TYPE_VIDEO\x10\x05\x12\x17\n\x13\x43OLUMN_TYPE_ECHARTS\x10\x06\x12\x18\n\x14\x43OLUMN_TYPE_OBJECT3D\x10\x07\x12\x18\n\x14\x43OLUMN_TYPE_MOLECULE\x10\x08\x12\x14\n\x10\x43OLUMN_TYPE_HTML\x10\t*]\n\x0b\x43olumnClass\x12\x1c\n\x18\x43OLUMN_CLASS_UNSPECIFIED\x10\x00\x12\x17\n\x13\x43OLUMN_CLASS_CUSTOM\x10\x01\x12\x17\n\x13\x43OLUMN_CLASS_SYSTEM\x10\x02*\x9d\x02\n\tChartType\x12\x1a\n\x16\x43HART_TYPE_UNSPECIFIED\x10\x00\x12\x13\n\x0f\x43HART_TYPE_LINE\x10\x01\x12\x12\n\x0e\x43HART_TYPE_BAR\x10\x02\x12\x15\n\x11\x43HART_TYPE_SCALAR\x10\x03\x12\x14\n\x10\x43HART_TYPE_IMAGE\x10\x04\x12\x14\n\x10\x43HART_TYPE_AUDIO\x10\x05\x12\x13\n\x0f\x43HART_TYPE_TEXT\x10\x06\x12\x14\n\x10\x43HART_TYPE_VIDEO\x10\x07\x12\x16\n\x12\x43HART_TYPE_ECHARTS\x10\x08\x12\x17\n\x13\x43HART_TYPE_OBJECT3D\x10\t\x12\x17\n\x13\x43HART_TYPE_MOLECULE\x10\n\x12\x13\n\x0f\x43HART_TYPE_HTML\x10\x0b*\x8f\x01\n\x0bSectionType\x12\x1c\n\x18SECTION_TYPE_UNSPECIFIED\x10\x00\x12\x17\n\x13SECTION_TYPE_PINNED\x10\x01\x12\x17\n\x13SECTION_TYPE_HIDDEN\x10\x02\x12\x17\n\x13SECTION_TYPE_PUBLIC\x10\x03\x12\x17\n\x13SECTION_TYPE_SYSTEM\x10\x04\x42JZHgithub.com/swanhubx/swanlab/core/proto/swanlab/metric/column/v1;columnv1b\x06proto3') _globals = globals() _builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals) @@ -32,18 +32,18 @@ if not _descriptor._USE_C_DESCRIPTORS: _globals['DESCRIPTOR']._loaded_options = None _globals['DESCRIPTOR']._serialized_options = b'ZHgithub.com/swanhubx/swanlab/core/proto/swanlab/metric/column/v1;columnv1' - _globals['_COLUMNTYPE']._serialized_start=669 - _globals['_COLUMNTYPE']._serialized_end=924 - _globals['_COLUMNCLASS']._serialized_start=926 - _globals['_COLUMNCLASS']._serialized_end=1019 - _globals['_CHARTTYPE']._serialized_start=1022 - _globals['_CHARTTYPE']._serialized_end=1307 - _globals['_SECTIONTYPE']._serialized_start=1310 - _globals['_SECTIONTYPE']._serialized_end=1453 + _globals['_COLUMNTYPE']._serialized_start=733 + _globals['_COLUMNTYPE']._serialized_end=988 + _globals['_COLUMNCLASS']._serialized_start=990 + _globals['_COLUMNCLASS']._serialized_end=1083 + _globals['_CHARTTYPE']._serialized_start=1086 + _globals['_CHARTTYPE']._serialized_end=1371 + _globals['_SECTIONTYPE']._serialized_start=1374 + _globals['_SECTIONTYPE']._serialized_end=1517 _globals['_YRANGE']._serialized_start=67 _globals['_YRANGE']._serialized_end=127 _globals['_METRICCOLORS']._serialized_start=129 _globals['_METRICCOLORS']._serialized_end=172 _globals['_COLUMNRECORD']._serialized_start=175 - _globals['_COLUMNRECORD']._serialized_end=666 + _globals['_COLUMNRECORD']._serialized_end=730 # @@protoc_insertion_point(module_scope) diff --git a/swanlab/proto/swanlab/metric/column/v1/column_pb2.pyi b/swanlab/proto/swanlab/metric/column/v1/column_pb2.pyi index d7d8389a1..ec92381db 100644 --- a/swanlab/proto/swanlab/metric/column/v1/column_pb2.pyi +++ b/swanlab/proto/swanlab/metric/column/v1/column_pb2.pyi @@ -95,7 +95,7 @@ class MetricColors(_message.Message): def __init__(self, light: _Optional[str] = ..., dark: _Optional[str] = ...) -> None: ... class ColumnRecord(_message.Message): - __slots__ = ("column_class", "column_type", "column_key", "column_name", "section_name", "section_type", "y_range", "chart_index", "chart_name", "chart_type", "metric_name", "metric_colors") + __slots__ = ("column_class", "column_type", "column_key", "column_name", "section_name", "section_type", "y_range", "chart_index", "chart_name", "chart_type", "metric_name", "metric_colors", "x_axis", "hidden") COLUMN_CLASS_FIELD_NUMBER: _ClassVar[int] COLUMN_TYPE_FIELD_NUMBER: _ClassVar[int] COLUMN_KEY_FIELD_NUMBER: _ClassVar[int] @@ -108,6 +108,8 @@ class ColumnRecord(_message.Message): CHART_TYPE_FIELD_NUMBER: _ClassVar[int] METRIC_NAME_FIELD_NUMBER: _ClassVar[int] METRIC_COLORS_FIELD_NUMBER: _ClassVar[int] + X_AXIS_FIELD_NUMBER: _ClassVar[int] + HIDDEN_FIELD_NUMBER: _ClassVar[int] column_class: ColumnClass column_type: ColumnType column_key: str @@ -120,4 +122,6 @@ class ColumnRecord(_message.Message): chart_type: ChartType metric_name: str metric_colors: MetricColors - def __init__(self, column_class: _Optional[_Union[ColumnClass, str]] = ..., column_type: _Optional[_Union[ColumnType, str]] = ..., column_key: _Optional[str] = ..., column_name: _Optional[str] = ..., section_name: _Optional[str] = ..., section_type: _Optional[_Union[SectionType, str]] = ..., y_range: _Optional[_Union[YRange, _Mapping]] = ..., chart_index: _Optional[str] = ..., chart_name: _Optional[str] = ..., chart_type: _Optional[_Union[ChartType, str]] = ..., metric_name: _Optional[str] = ..., metric_colors: _Optional[_Union[MetricColors, _Mapping]] = ...) -> None: ... + x_axis: str + hidden: bool + def __init__(self, column_class: _Optional[_Union[ColumnClass, str]] = ..., column_type: _Optional[_Union[ColumnType, str]] = ..., column_key: _Optional[str] = ..., column_name: _Optional[str] = ..., section_name: _Optional[str] = ..., section_type: _Optional[_Union[SectionType, str]] = ..., y_range: _Optional[_Union[YRange, _Mapping]] = ..., chart_index: _Optional[str] = ..., chart_name: _Optional[str] = ..., chart_type: _Optional[_Union[ChartType, str]] = ..., metric_name: _Optional[str] = ..., metric_colors: _Optional[_Union[MetricColors, _Mapping]] = ..., x_axis: _Optional[str] = ..., hidden: bool = ...) -> None: ... diff --git a/swanlab/sdk/__init__.py b/swanlab/sdk/__init__.py index 9240e7aaf..0a15876a4 100644 --- a/swanlab/sdk/__init__.py +++ b/swanlab/sdk/__init__.py @@ -15,7 +15,7 @@ from .cmd.require import require from .cmd.run import ( async_log, - define_scalar, + define_metric, finish, log, log_audio, @@ -51,7 +51,7 @@ "log_object3d", "log_molecule", "async_log", - "define_scalar", + "define_metric", "save", "sync", "merge_settings", diff --git a/swanlab/sdk/cmd/run.py b/swanlab/sdk/cmd/run.py index 8c8a28aa9..023aec178 100644 --- a/swanlab/sdk/cmd/run.py +++ b/swanlab/sdk/cmd/run.py @@ -45,6 +45,6 @@ def wrapper(*args: Any, **kwargs: Any) -> Any: log_molecule = _make_run_cmd("log_molecule") log_html = _make_run_cmd("log_html") async_log = _make_run_cmd("async_log") -define_scalar = _make_run_cmd("define_scalar") +define_metric = _make_run_cmd("define_metric") finish = _make_run_cmd("finish") save = _make_run_cmd("save") diff --git a/swanlab/sdk/internal/bus/__init__.py b/swanlab/sdk/internal/bus/__init__.py index 0b2733a43..fcef27cd4 100644 --- a/swanlab/sdk/internal/bus/__init__.py +++ b/swanlab/sdk/internal/bus/__init__.py @@ -10,10 +10,11 @@ from .events import ( ConfigEvent, EventPayload, + FileSaveEvent, LogEvent, + MetricDefineEvent, MetricLogEvent, ParseResult, - ScalarDefineEvent, ) __all__ = [ @@ -21,9 +22,10 @@ "RunQueue", "EmitterProtocol", "MetricLogEvent", - "ScalarDefineEvent", + "MetricDefineEvent", "ConfigEvent", "LogEvent", + "FileSaveEvent", "EventPayload", "ParseResult", ] diff --git a/swanlab/sdk/internal/bus/events.py b/swanlab/sdk/internal/bus/events.py index 7a9d322f3..f4bc68e82 100644 --- a/swanlab/sdk/internal/bus/events.py +++ b/swanlab/sdk/internal/bus/events.py @@ -14,7 +14,6 @@ from swanlab.proto.swanlab.metric.data.v1.data_pb2 import MediaRecord, ScalarRecord from swanlab.proto.swanlab.terminal.v1.log_pb2 import LogLevel from swanlab.sdk.internal.context.transformer import TransformData -from swanlab.sdk.typings.run.column import ScalarXAxisType @dataclass @@ -27,23 +26,31 @@ class MetricLogEvent: @dataclass -class ScalarDefineEvent: - """显式创建标量列事件""" +class MetricDefineEvent: + """define_metric 事件,携带参数归一化后的定义补丁。 - # 指标键 + 字段为两态:``None`` 表示 "未提供",具体值表示使用该值: + + - 公开签名无法区分 "省略" 与 "显式 None"(``x_axis`` / ``section_name``), + 两者均为 ``None``("未提供");merge 时保留旧值,replace 时重置为默认。 + 清除旧值必须使用 ``overwrite=True``。 + - ``hidden`` / ``step_sync`` 的 ``False`` 是有效值,不会被收成 "未提供"。 + """ + + # rule identity(exact key 或 glob pattern) key: str - # 指标名 - name: Optional[str] = None - # 指标颜色 - color: Optional[str] = None - # 是否为系统指标 - system: bool = False - # x轴,可以是其他的标量,也可以是系统值"_step"或"_relative_time" - x_axis: Optional[ScalarXAxisType] = None - # 图表索引 - chart: Optional[str] = None - # 图表名 - chart_name: Optional[str] = None + # key 是否为末尾单个 '*' 的 glob 模式;由上层校验分类后传入,resolver 不再解析字符串 + is_glob: bool = False + # None=not provided, str=custom X key + x_axis: Optional[str] = None + # None=not provided, str=named section + section_name: Optional[str] = None + # None=not provided, True/False + hidden: Optional[bool] = None + # None=not provided, True/False。仅对 custom X 有意义:该 Y 是否触发跨 step 的 X 注入 + step_sync: Optional[bool] = None + # merge(False)vs replace(True) + overwrite: bool = False @dataclass @@ -75,7 +82,7 @@ class FileSaveEvent: # 事件载体类型 EventPayload = Union[ MetricLogEvent, - ScalarDefineEvent, + MetricDefineEvent, ConfigEvent, LogEvent, FileSaveEvent, diff --git a/swanlab/sdk/internal/core_python/transport/sender.py b/swanlab/sdk/internal/core_python/transport/sender.py index d32a20fba..797c3da6a 100644 --- a/swanlab/sdk/internal/core_python/transport/sender.py +++ b/swanlab/sdk/internal/core_python/transport/sender.py @@ -148,11 +148,7 @@ def upload(self, record_type: str, records: Sequence[Record]) -> None: self._tracker.advance_records(len(records)) def upload_column(self, records: Sequence[Record], batch_size: int = 3000) -> None: - columns = [] - for record in records: - r = encode_column(record.column) - if r: - columns.append(r) + columns = [encode_column(record.column) for record in records] if not columns: return for i in range(0, len(columns), batch_size): @@ -563,7 +559,7 @@ def on_read(current_bytes): return completed_parts -def encode_column(record: ColumnRecord) -> Optional[UploadColumn]: +def encode_column(record: ColumnRecord) -> UploadColumn: """ 将列记录编码为后端所需的格式(DTO) """ @@ -571,8 +567,12 @@ def encode_column(record: ColumnRecord) -> Optional[UploadColumn]: column: UploadColumn = {"key": record.column_key, "type": adapter.column[record.column_type]} if record.column_class == ColumnClass.COLUMN_CLASS_SYSTEM: column["type"] = "SYSTEM" + if record.HasField("x_axis"): + column["xAxis"] = record.x_axis if record.section_name: column["sectionName"] = record.section_name + if record.hidden: + column["hidden"] = True return column diff --git a/swanlab/sdk/internal/pkg/constraints/__init__.py b/swanlab/sdk/internal/pkg/constraints/__init__.py index 5690101e2..e1884867e 100644 --- a/swanlab/sdk/internal/pkg/constraints/__init__.py +++ b/swanlab/sdk/internal/pkg/constraints/__init__.py @@ -117,8 +117,8 @@ def _no_dot_slash_edges(v: str) -> str: ] """Metric / log key: 1-512 chars, no control characters, must not start or end with ``'.'`` or ``'/'``.""" -MetricName = Annotated[str, Field(max_length=512)] -"""Display label for metrics and runs: up to 512 chars.""" +MetricName = Annotated[str, Field(max_length=512, pattern=r"^[^\x00-\x1f\x7f]+$")] +"""Display label for metrics and runs: up to 512 chars, no control characters.""" ChartName = Annotated[str, Field(min_length=1, max_length=512)] """Display name for charts: 1-512 chars.""" diff --git a/swanlab/sdk/internal/pkg/helper/__init__.py b/swanlab/sdk/internal/pkg/helper/__init__.py index 4e54c1429..81b465430 100644 --- a/swanlab/sdk/internal/pkg/helper/__init__.py +++ b/swanlab/sdk/internal/pkg/helper/__init__.py @@ -9,7 +9,7 @@ from .env import DEBUG, is_interactive, is_jupyter from .impl import apply_requirement, get_core_impl, get_probe_impl, set_core_impl, set_probe_impl -from .system import fmt_system_key, is_system_key +from .metric import fmt_system_key, is_custom_x, is_system_key from .version import get_swanlab_latest_version, get_swanlab_version __all__ = [ @@ -17,6 +17,7 @@ "is_jupyter", "is_interactive", "is_system_key", + "is_custom_x", "strip_none", "fmt_run_path", "fmt_system_key", diff --git a/swanlab/sdk/internal/pkg/helper/system.py b/swanlab/sdk/internal/pkg/helper/metric.py similarity index 50% rename from swanlab/sdk/internal/pkg/helper/system.py rename to swanlab/sdk/internal/pkg/helper/metric.py index 5c5fe6f7d..1eedb565f 100644 --- a/swanlab/sdk/internal/pkg/helper/system.py +++ b/swanlab/sdk/internal/pkg/helper/metric.py @@ -1,14 +1,17 @@ """ @author: cunyue -@file: system.py +@file: metric.py @time: 2026/5/5 19:31 -@description: 系统指标 +@description: 指标 key 工具(系统指标 key 前缀、系统 X 轴的判定) """ -__all__ = ["fmt_system_key", "is_system_key"] +__all__ = ["fmt_system_key", "is_system_key", "is_custom_x"] _SYSTEM_KEY_PREFIX = "__swanlab__." +# 系统 X 轴:由服务端/SDK 内部维护,不视为自定义 metric key +_SYSTEM_X_AXES = frozenset({"_step", "_relative_time"}) + def fmt_system_key(key: str): """ @@ -26,3 +29,12 @@ def is_system_key(key: str): :return: 是否是系统指标key """ return key.startswith(_SYSTEM_KEY_PREFIX) + + +def is_custom_x(x_axis: str) -> bool: + """ + 判断 x_axis 是否为自定义(非系统 _step / _relative_time)的 scalar key + :param x_axis: X 轴 key + :return: 是否为自定义 X 轴 + """ + return x_axis not in _SYSTEM_X_AXES diff --git a/swanlab/sdk/internal/probe_python/typings/__init__.py b/swanlab/sdk/internal/probe_python/typings/__init__.py index 1b255b132..765aa6d75 100644 --- a/swanlab/sdk/internal/probe_python/typings/__init__.py +++ b/swanlab/sdk/internal/probe_python/typings/__init__.py @@ -457,8 +457,8 @@ class SystemEnvironment(BaseModel): class SystemScalar(BaseModel): """系统监控标量信息,作为定义标量中间载体 - 与 run.define_scalar() 参数对齐,由系统监控模块内部使用, - 用于在硬件监控线程启动前批量注册系统标量定义。 + 由系统监控模块内部使用,用于在硬件监控线程启动前 + 批量注册系统标量定义。 """ key: str = Field(..., pattern=r"^[a-z0-9\.\-]+$", max_length=512, min_length=1) diff --git a/swanlab/sdk/internal/run/__init__.py b/swanlab/sdk/internal/run/__init__.py index 38dbde2df..510ade84f 100644 --- a/swanlab/sdk/internal/run/__init__.py +++ b/swanlab/sdk/internal/run/__init__.py @@ -27,8 +27,7 @@ ) from swanlab.proto.swanlab.grpc.probe.v1.probe_pb2 import DeliverProbeStartRequest from swanlab.proto.swanlab.run.v1.run_pb2 import FinishRecord -from swanlab.sdk.internal.bus import MetricLogEvent -from swanlab.sdk.internal.bus.events import FileSaveEvent +from swanlab.sdk.internal.bus import FileSaveEvent, MetricDefineEvent, MetricLogEvent from swanlab.sdk.internal.context import RunContext from swanlab.sdk.internal.core_python import client from swanlab.sdk.internal.pkg import adapter, console, fork, helper, safe @@ -48,7 +47,6 @@ normalize_media_input, ) from swanlab.sdk.typings.run import AsyncLogType, FinishType, ModeType, SaveType -from swanlab.sdk.typings.run.column import ScalarXAxisType from swanlab.sdk.typings.run.transforms import CaptionsType from swanlab.sdk.typings.run.transforms.audio import AudioDatasType, AudioRatesType from swanlab.sdk.typings.run.transforms.echarts import EChartsDatasType @@ -611,59 +609,153 @@ def log_html(self, *, key: str, data: HtmlDatasType, caption: CaptionsType = Non normalized_data = normalize_media_input(Html, data, caption=caption) self.log({key: normalized_data}, step=step) - @with_api("run.define_scalar()") - def define_scalar( + @with_api("run.define_metric()") + def define_metric( self, - *, key: str, - name: Optional[str] = None, - color: Optional[str] = None, - x_axis: Optional[ScalarXAxisType] = None, - chart_name: Optional[str] = None, + *, + x_axis: Optional[str] = None, + section_name: Optional[str] = None, + hidden: Optional[bool] = None, + step_sync: Optional[bool] = None, + overwrite: bool = False, + **kwargs, ): - """ - Manually define a scalar column before logging. + """Define a metric's display configuration before logging. + + Customizes how an auto-generated chart for ``key`` appears in project Views: + X-axis, section placement, and visibility. The same ``(class, key)`` + shares one chart across all runs in the project. + + :param key: Metric key. Supports exact match and a single trailing ``*`` + glob (e.g. ``"train/*"``). System keys are never matched. + :param x_axis: Custom X-axis key. ``None`` (default) means the system + step. The system axes ``"_step"`` and ``"_relative_time"`` are also + accepted: they are resolved by the server, and the SDK performs no + step injection or X dedup for them. Only affects scalar charts; + media ignores this. The X series is assumed monotonically + non-decreasing — for a given X value, only the first logged Y is + kept (consecutive-duplicate X values are dropped). ``step_metric`` + is accepted as a ``**kwargs`` alias for backward compatibility. + :param section_name: Section name for the auto chart. ``None`` means + the default section derived from the key. + :param hidden: If ``True``, place the chart in the HIDDEN section. + Three states: ``None`` (default) means "not provided" — merge mode + keeps the previous value; ``True`` hides; ``False`` explicitly + unhides (also effective in merge mode). + :param step_sync: Whether this Y key should copy the latest custom X + value onto the current step when X and Y are logged separately. + Defaults to ``True`` when ``x_axis`` (or ``step_metric``) is set, + ``False`` otherwise — without a custom X axis it has no effect. + ``False`` means this Y will not trigger X injection: the two series + only align when they share a step (same ``log()`` or the same + explicit ``step``). Duplicate-X dropping then uses only an X value + present in this event. Sibling Y keys that keep ``step_sync=True`` + can still inject X for the shared series. + :param overwrite: If ``False`` (default), merge with previous calls for + the same ``key`` — unspecified fields reuse the previous value. + If ``True``, unspecified fields reset to their default, overwriting + previous values. Only affects rules not yet applied to a logged key. + + .. note:: + Project-wide first-definition-effective. A chart is shared across every run + in the project under the same ``(class, key)``, and + ``define_metric`` only shapes it on the **first** run that + introduces the key to the project. Once the chart exists — created + by this run's first log or by any earlier run — later + ``define_metric`` calls (including ``overwrite=True``) have no + effect, and no warning is raised. To change an existing chart, + edit it in the UI. - :param key: The key for the scalar column. Supports wildcards (e.g. ``"train/*"``) to match multiple columns. - :param name: Optional display name for the column. - :param color: Optional color for the column, as a hex color code. - :param x_axis: Optional x-axis type for the column. - :param chart_name: Optional chart name to group the column into. - """ - raise NotImplementedError("run.define_scalar() is not available yet. Support is planned for a future release.") + Examples: - # TODO: 实现 glob 匹配逻辑 - # if not (this_key := fmt.safe_validate_key(key)): - # return console.error( - # f"Invalid key for define scalar: {key}, please use valid characters (alphanumeric, '.', '-', '/') and avoid special characters." - # ) - # - # original_name = name - # if name and not (name := fmt.safe_validate_name(name)): - # return console.error(f"Invalid name for define scalar: {original_name}, must be a string.") - # - # original_color = color - # if color and not (color := fmt.safe_validate_color(color)): - # return console.error(f"Invalid color for define scalar: {original_color}, must be a hex color code.") - # - # if (this_x_axis := fmt.safe_validate_x_axis(x_axis)) is None: - # return console.error(f"Invalid x_axis for define scalar: {x_axis}, must be a valid ScalarXAxisType.") - # - # original_chart_name = chart_name - # if chart_name and not (chart_name := fmt.safe_validate_chart_name(chart_name)): - # return console.error(f"Invalid chart_name for define scalar: {original_chart_name}, must be a string.") + Custom X-axis with separate logging: + + >>> import swanlab + >>> swanlab.init(mode="local") + >>> run = swanlab.get_run() + >>> run.define_metric("train/loss", x_axis="train/epoch") + >>> swanlab.log({"train/epoch": 1}) # step=1 + >>> swanlab.log({"train/loss": 0.8}) # step=2, auto-syncs train/epoch=1 + + Glob pattern for all validation metrics: + + >>> run.define_metric("val/*", section_name="Validation") + + Disable X injection (X/Y must share a step to align): + + >>> run.define_metric("train/acc", x_axis="train/epoch", step_sync=False) + >>> swanlab.log({"train/epoch": 1}, step=1) + >>> swanlab.log({"train/acc": 0.9}, step=2) # no injected epoch at step=2 + + ``step_metric`` alias: + + >>> run.define_metric("loss", step_metric="custom_step") + """ + # 1. key 校验,完全复用 log 侧的 validate_key(非字符串强转 str、清洗首尾空白与'./'、超长截断) + # 保证规则与 log 实际产生的规范 key 落在同一形式,避免发出永不匹配的规则 # - # self._components.emitter.emit( - # ScalarDefineEvent( - # key=this_key, - # name=name, - # color=color, - # system=False, - # x_axis=this_x_axis, - # chart_name=chart_name, - # chart=None, - # ) - # ) + # 边界情况:>512 的 key 截断可能截掉末尾 glob '*',glob 退化为 exact 匹配,此时属病态输入且 log 侧同样截断、两边落点一致,暂不特殊处理 + try: + this_key = fmt.validate_key(key) + except ValueError as e: + console.error(f"Invalid key for define_metric: {key!r}, {e}") + return + + # 2. glob 校验与分类:仅支持末尾单个 '*'(如 train/*),拒绝 *loss、train/*/loss、train/** 等 + star_count = this_key.count("*") + if star_count > 0 and not (star_count == 1 and this_key.endswith("*")): + console.error( + f"Invalid glob pattern for define_metric: {key!r}, " + f"only a single trailing '*' is supported (e.g. 'train/*')" + ) + return + is_glob = star_count == 1 + + # 3. step_metric 兼容别名(仅从 kwargs 读取;其余未知 kwargs 静默忽略,与 init 兼容风格一致) + step_metric = kwargs.pop("step_metric", None) + if step_metric is not None: + if x_axis is None: + x_axis = step_metric + elif x_axis != step_metric: + console.warning( + f"Conflicting x_axis={x_axis!r} and step_metric={step_metric!r} " + f"for key {this_key!r}, using x_axis={x_axis!r}" + ) + + # 4. x_axis / section_name 校验 + if x_axis is not None: + validated_x = fmt.safe_validate_x_axis(x_axis) + if validated_x is None: + console.error( + f"Invalid x_axis for define_metric: {x_axis}, must be a valid metric key, " + f"'_step', '_relative_time', and must not be a system metric key." + ) + return + x_axis = validated_x + # section_name:非法 / 空串降级为默认 section(不阻断整条 define) + if section_name is not None: + validated_section = fmt.safe_validate_name(section_name) + if validated_section is None or not validated_section.strip(): + console.warning( + f"Invalid section_name for define_metric: {section_name!r}, ignored; using default section." + ) + section_name = None + else: + section_name = validated_section + + # 5. 发射事件(三态字段校验后即为 None|str / None|bool,直接透传;is_glob 已在第 2 步分类) + self._components.emitter.emit( + MetricDefineEvent( + key=this_key, + is_glob=is_glob, + x_axis=x_axis, + section_name=section_name, + hidden=hidden, + step_sync=step_sync, + overwrite=overwrite, + ) + ) @with_api("run.save()") def save( diff --git a/swanlab/sdk/internal/run/components/__init__.py b/swanlab/sdk/internal/run/components/__init__.py index 504589843..3c24aad7d 100644 --- a/swanlab/sdk/internal/run/components/__init__.py +++ b/swanlab/sdk/internal/run/components/__init__.py @@ -11,7 +11,6 @@ from swanlab.sdk.internal.context import RunContext from swanlab.sdk.internal.pkg import console, fork from swanlab.sdk.internal.run.components.asynctask import AsyncTaskManager -from swanlab.sdk.internal.run.components.builder import RecordBuilder from swanlab.sdk.internal.run.components.config import ( Config, create_run_config, @@ -28,7 +27,6 @@ "ConsumerProtocol", "TerminalProxyProtocol", "AsyncTaskManager", - "RecordBuilder", "BackgroundConsumer", "NullConsumer", "NullEmitter", @@ -53,9 +51,8 @@ def __init__(self, ctx: RunContext) -> None: # 核心组件 self._asynctask = AsyncTaskManager() - self._builder = RecordBuilder(ctx) self._emitter = _factory_emitter(ctx) - self._consumer = _factory_consumer(ctx, self._emitter, self._builder) + self._consumer = _factory_consumer(ctx, self._emitter) self._config = _factory_config(ctx, self._emitter) self._terminal: TerminalProxyProtocol = _factory_terminal(ctx, self._emitter, self._init_pid) @@ -113,8 +110,8 @@ def stop(self, async_log_timeout: float | None = None) -> None: # 3. 解绑 config deactivate_run_config() - # 4. 停止消费者线程(消费剩余事件包括最后的 ConsoleEvent) - console.debug("SwanLab Run is finishing, waiting for logs to flush...") + # 4. 停止消费者线程 + console.debug("Swanlab Run is finishing, waiting for logs to flush...") self._consumer.stop() self._consumer.join() @@ -133,11 +130,10 @@ def _factory_emitter(ctx: RunContext) -> EmitterProtocol: def _factory_consumer( ctx: RunContext, e: EmitterProtocol, - b: RecordBuilder, ) -> ConsumerProtocol: if ctx.config.settings.mode == "disabled": - return NullConsumer(ctx, e.queue, b) - return BackgroundConsumer(ctx, e.queue, b) + return NullConsumer(ctx, e.queue) + return BackgroundConsumer(ctx, e.queue) def _factory_config(ctx: RunContext, e: EmitterProtocol) -> Config: diff --git a/swanlab/sdk/internal/run/components/consumer/__init__.py b/swanlab/sdk/internal/run/components/consumer/__init__.py index ef2a6dad9..94cbc6ef9 100644 --- a/swanlab/sdk/internal/run/components/consumer/__init__.py +++ b/swanlab/sdk/internal/run/components/consumer/__init__.py @@ -5,13 +5,14 @@ @description: 后台消费者组件,用于消费运行事件 """ +import math import queue import threading from abc import ABC -from typing import TYPE_CHECKING, List, Tuple +from typing import Any, Dict, List, Tuple, Union from swanlab.proto.swanlab.metric.column.v1.column_pb2 import ColumnRecord -from swanlab.proto.swanlab.metric.data.v1.data_pb2 import MediaRecord, ScalarRecord +from swanlab.proto.swanlab.metric.data.v1.data_pb2 import MediaRecord, ScalarRecord, ScalarValue from swanlab.proto.swanlab.save.v1.save_pb2 import SaveRecord from swanlab.proto.swanlab.terminal.v1.log_pb2 import LogRecord from swanlab.sdk.internal.bus.emitter import RunQueue @@ -19,15 +20,15 @@ ConfigEvent, FileSaveEvent, LogEvent, + MetricDefineEvent, MetricLogEvent, - ScalarDefineEvent, ) from swanlab.sdk.internal.context import RunContext -from swanlab.sdk.internal.pkg import console, safe - -if TYPE_CHECKING: - from swanlab.sdk.internal.run.components.builder import RecordBuilder +from swanlab.sdk.internal.pkg import console, helper, safe +from swanlab.sdk.internal.run.transforms import Scalar +from .builder import RecordBuilder, is_scalar_value +from .resolver import DefinitionResolver LogBatch = List[LogRecord] ColumnBatch = List[ColumnRecord] @@ -44,11 +45,10 @@ def __init__( self, ctx: RunContext, event_queue: RunQueue, - builder: "RecordBuilder", flush_timeout: float = 0.5, batch_size: int = 100, ): - _ = (ctx, event_queue, builder, flush_timeout, batch_size) + _ = (ctx, event_queue, flush_timeout, batch_size) def start(self) -> None: ... @@ -69,14 +69,15 @@ def __init__( self, ctx: RunContext, event_queue: RunQueue, - builder: "RecordBuilder", flush_timeout: float = 0.5, batch_size: int = 100, ): - super().__init__(ctx, event_queue, builder, flush_timeout, batch_size) + super().__init__(ctx, event_queue, flush_timeout, batch_size) self._ctx = ctx self._queue = event_queue - self._builder = builder + # builder / resolver 均由 consumer 独占创建与维护,所有状态访问都发生在本线程内 + self._builder = RecordBuilder(ctx) + self._resolver = DefinitionResolver() self._core = ctx.core self._flush_timeout = flush_timeout self._batch_size = batch_size @@ -169,12 +170,14 @@ def _run(self) -> None: self._handle_event(event) if self._batch_full: self._flush() + # 线程退出时清理状态 + self._resolver.clear() def _handle_event(self, event) -> None: if isinstance(event, MetricLogEvent): self._handle_metric_log(event) - elif isinstance(event, ScalarDefineEvent): - self._handle_scalar_define(event) + elif isinstance(event, MetricDefineEvent): + self._resolver.handle_define(event) elif isinstance(event, ConfigEvent): self._save_batch.append(self._builder.build_config(event)) elif isinstance(event, LogEvent): @@ -184,20 +187,161 @@ def _handle_event(self, event) -> None: self._save_batch.append(self._builder.build_save(event)) def _handle_metric_log(self, event: MetricLogEvent) -> None: - for key, value in event.data.items(): + """处理指标日志事件:构建 scalar/media 记录;调用过 define_metric 时额外物化列定义, + 并为 custom X 轴指标做跨 step 的 X/Y 对齐。 + + 按是否登记过 define rule 分流: + - 未调用过 define_metric → _log_plain:纯构建,列由 core 收到数据后自动创建; + - 存在 define rule → _log_define:物化 define 过的列 + custom X 对齐机制。 + """ + # 1. 过滤用户伪造的系统前缀 key + raw_data = event.data + data: Dict[str, Any] = {} + for key, value in raw_data.items(): + if helper.is_system_key(key): + # 系统列名不允许用户伪造,避免覆盖 swanlab 内部列定义 + # 这种情况比较少见,并且通常为有意为之,因此打印warning而非debug + console.warning(f"Metric '{key}' at step {event.step} is a system key, skipped") + continue + data[key] = value + + if self._resolver.has_rules: + self._log_define(event, data) + else: + self._log_plain(event, data) + + def _log_plain(self, event: MetricLogEvent, data: Dict[str, Any]) -> None: + """无 define 路径:纯构建 data record,不产出任何 ColumnRecord——列由 core + 收到数据后自动创建(与未引入 define_metric 时的行为一致)。 + + 仍调用 resolve_concrete 登记 automatic 状态:钉住 "key 首次 log 后 define + 不再生效" 的契约,之后到达的 define 无法认领已 log 过的 key。 + """ + for key, value in data.items(): with safe.block(message=f"Error when parsing metric '{key}'"): data_record, _ = self._builder.build_scalar_or_media(value, key, event.timestamp, event.step) if data_record is None: console.warning(f"Metric '{key}' at step {event.step} returned no data, skipped") continue - if isinstance(data_record, ScalarRecord): - self._scalar_batch.append(data_record) + metric_class = "SCALAR" if isinstance(data_record, ScalarRecord) else "MEDIA" + self._resolver.resolve_concrete(key, metric_class) # 仅登记钉住,不物化 + self._append_data_record(data_record) + + def _log_define(self, event: MetricLogEvent, data: Dict[str, Any]) -> None: + """有 define 路径:预扫描/候选注入(存在 custom X 规则时)→ 构建 + 物化 define + 过的列 + custom X 的 Y 按 X 值去重 → 提交注入。 + """ + # 2. 预扫描,仅在登记过 custom X 规则时执行:收集显式 scalar 值并更新 custom X 缓存 + explicit_scalars: dict[str, float] = {} + scalar_values: dict[str, ScalarValue] = {} + if self._resolver.has_custom_x: + for key, value in data.items(): + if not is_scalar_value(value): + continue + # 预扫描失败不单独报 error:第 3 步 build 时同一值会再次 transform, + # 届时以完整上下文上报;此处仅留 debug 级别 + with safe.block(message=f"Error when pre-scanning scalar '{key}'", level="debug", write_to_tty=False): + scalar_value = Scalar.transform(value) + # transform 结果(含 nan)缓存供第 4 步复用,避免 build 时二次提取 tensor/numpy + scalar_values[key] = scalar_value + val = scalar_value.number + if math.isfinite(val): + explicit_scalars[key] = val + # 收集完毕,更新 custom X 缓存,此时搜集仅显式标量值,避免 media/nan 注入 + for key, val in explicit_scalars.items(): + self._resolver.update_custom_x_cache(key, val, event.step) + + # 3. 计算候选注入,为 custom X 且 event 未含 X 的 Y 生成最近真实 X 值 + # candidate_x 将是最终上报的注入 X 值 + candidate_x: dict[str, float] = {} + for key in explicit_scalars: + concrete = self._resolver.resolve_concrete(key, "SCALAR") + # 3.1 显式 step_sync=False 的 Y 不参与注入 + if not concrete.effective.step_sync: + continue + # 3.2 如果x轴为系统内部step,则不参与注入 + x_axis = concrete.effective.x_axis + if not helper.is_custom_x(x_axis): + continue + # 3.3 如果用户在本次 event 里显式 log 了 X 值,则不注入 + if x_axis in data or x_axis in candidate_x: + continue + # 3.4 从 cache 取最近真实 X 值 + cached_x = self._resolver.get_custom_x(x_axis) + if cached_x is None: + continue + candidate_x[x_axis] = cached_x + + # 4. 把本次 log 的每个 key 变成 record,包括: + # - 首次物化 ColumnRecord + # - 构建 ScalarRecord / MediaRecord + # - custom X 标量按 X 值去重,记录被存活 Y 消费的候选 X + consumed_candidate_x: set[str] = set() + for key, value in data.items(): + with safe.block(message=f"Error when parsing metric '{key}'"): + # 4.1 复用预扫描的 transform 结果构建 record,避免对 tensor/numpy 的二次 .item() 提取 + cached = scalar_values.get(key) + if cached is not None: + data_record = Scalar.build_data_record( + key=key, step=event.step, timestamp=event.timestamp, data=cached + ) else: - self._media_batch.append(data_record) + data_record, _ = self._builder.build_scalar_or_media(value, key, event.timestamp, event.step) + if data_record is None: + console.warning(f"Metric '{key}' at step {event.step} returned no data, skipped") + continue - def _handle_scalar_define(self, event: ScalarDefineEvent) -> None: - this_column = self._builder.build_column_from_scalar_define(event) - self._column_batch.append(this_column) + # 4.2 物化 ColumnRecord,resolver侧保证仅对define过的 key 物化一次,后续 log 不再重复发 ColumnRecord + # 对于未define的key,物化行为依旧交给core自动处理 + is_scalar = isinstance(data_record, ScalarRecord) + metric_class = "SCALAR" if is_scalar else "MEDIA" + concrete = self._resolver.resolve_concrete(key, metric_class) + col = self._resolver.materialize_column(key, metric_class, data_record.type) + if col is not None: + self._column_batch.append(col) + + # 4.3 custom X 轴 first-writer-wins:同一 X 值上首次 Y 值为准 + x_axis = concrete.effective.x_axis + if is_scalar and helper.is_custom_x(x_axis): + x_value = explicit_scalars.get(x_axis) + used_candidate = False + # 进行 step 对齐,如果当前log没有显式 log X 值,则尝试从候选注入中取最近的真实 X 值 + if x_value is None and concrete.effective.step_sync: + x_value = candidate_x.get(x_axis) + used_candidate = x_value is not None + # 对本次提交的(x, y)进行x轴值的去重,比如(1,2)、(1,3),仅保留第一条(1,2),第二条(1,3)会被丢弃 + if x_value is not None and not self._resolver.try_accept_x_value(key, x_value): + console.debug( + f"Skipping '{key}': duplicate X value {x_value} for '{x_axis}'", + write_to_tty=False, + write_to_file=True, + ) + continue + if used_candidate: + consumed_candidate_x.add(x_axis) + + self._append_data_record(data_record) + + # 5. 提交注入X值:只有当某个 Y 真正用上了缓存 X 值、且该 Y 没被去重丢弃时,才为这个 X 构建 record 并标记本 step 已注入。 + for x_axis in candidate_x: + if x_axis not in consumed_candidate_x: + continue + if not self._resolver.try_inject_x(x_axis, event.step): + continue + with safe.block(message=f"Error when injecting custom X '{x_axis}'"): + injected_record, _ = self._builder.build_scalar_or_media( + candidate_x[x_axis], x_axis, event.timestamp, event.step + ) + # 注入值恒为标量(来自 cache 的 float);isinstance 同时收窄 pyright 类型 + if isinstance(injected_record, ScalarRecord): + self._scalar_batch.append(injected_record) + + def _append_data_record(self, data_record: Union[ScalarRecord, MediaRecord]) -> None: + """按类型把 data record 追加到对应批次。""" + if isinstance(data_record, ScalarRecord): + self._scalar_batch.append(data_record) + else: + self._media_batch.append(data_record) def _flush(self) -> None: if self._batch_empty: @@ -218,7 +362,13 @@ def _flush(self) -> None: log_batch = [] if column_batch: - self._core.upsert_columns(column_batch) + # 一次 flush 内按 (key, type) coalesce,保留最后一条 + coalesced: dict[tuple[str, int], ColumnRecord] = {} + for col in column_batch: + coalesced[(col.column_key, col.column_type)] = col + coalesced_list = list(coalesced.values()) + console.debug(f"Flushing column records to core: count={len(coalesced_list)}", write_to_tty=False) + self._core.upsert_columns(coalesced_list) column_batch = [] if scalar_batch: diff --git a/swanlab/sdk/internal/run/components/builder/__init__.py b/swanlab/sdk/internal/run/components/consumer/builder/__init__.py similarity index 85% rename from swanlab/sdk/internal/run/components/builder/__init__.py rename to swanlab/sdk/internal/run/components/consumer/builder/__init__.py index 236dfedba..adba64804 100644 --- a/swanlab/sdk/internal/run/components/builder/__init__.py +++ b/swanlab/sdk/internal/run/components/consumer/builder/__init__.py @@ -10,23 +10,29 @@ from google.protobuf.timestamp_pb2 import Timestamp -from swanlab.proto.swanlab.metric.column.v1.column_pb2 import ( - ColumnClass, - ColumnRecord, - ColumnType, - MetricColors, - SectionType, -) from swanlab.proto.swanlab.metric.data.v1.data_pb2 import MediaRecord from swanlab.proto.swanlab.save.v1.save_pb2 import SaveRecord, SaveType from swanlab.proto.swanlab.terminal.v1.log_pb2 import LogRecord -from swanlab.sdk.internal.bus.events import ConfigEvent, FileSaveEvent, LogEvent, ParseResult, ScalarDefineEvent +from swanlab.sdk.internal.bus.events import ConfigEvent, FileSaveEvent, LogEvent, ParseResult from swanlab.sdk.internal.context import RunContext, TransformMedia from swanlab.sdk.internal.pkg import adapter, console, fs from swanlab.sdk.internal.run.transforms import ECharts, Scalar, echarts _EchartsType = (echarts.Base, echarts.Table) +# 会被 build_scalar_or_media 处理为媒体(非标量)的类型 +_NON_SCALAR_TYPES = (list, TransformMedia, *_EchartsType) + + +def is_scalar_value(value: object) -> bool: + """判断是否会被 build_scalar_or_media 处理为标量 + + 在不触发 media transform 落盘的前提下决定是否预 transform。 + 判定必须与 build_scalar_or_media 的 dispatch 结果一致:显式注册类型走媒体 + 分支,默认分支中的 echarts 类型会被包装为 ECharts 媒体,其余按标量处理。 + """ + return not isinstance(value, _NON_SCALAR_TYPES) + class RecordBuilder: _MEDIA_MAX_SIZE = 10 * 1024**2 # 10 MB @@ -123,23 +129,6 @@ def _(self, value: TransformMedia, key: str, timestamp: Timestamp, step: int) -> ) return media_record, cls - @staticmethod - def build_column_from_scalar_define(event: ScalarDefineEvent) -> ColumnRecord: - """显式创建标量列(DefineEvent)""" - section_type = SectionType.SECTION_TYPE_SYSTEM if event.system else SectionType.SECTION_TYPE_PUBLIC - col = ColumnRecord( - column_key=event.key, - column_type=ColumnType.COLUMN_TYPE_SCALAR, - column_class=ColumnClass.COLUMN_CLASS_CUSTOM, - section_name=event.chart_name or "", - section_type=section_type, - chart_index=event.chart or "", - chart_name=event.chart_name or "", - metric_name=event.name or "", - metric_colors=MetricColors(light=event.color, dark=event.color) if event.color else None, - ) - return col - # ── 系统元数据 ── @staticmethod def build_config(event: ConfigEvent) -> SaveRecord: diff --git a/swanlab/sdk/internal/run/components/consumer/resolver/__init__.py b/swanlab/sdk/internal/run/components/consumer/resolver/__init__.py new file mode 100644 index 000000000..1351b4f9e --- /dev/null +++ b/swanlab/sdk/internal/run/components/consumer/resolver/__init__.py @@ -0,0 +1,342 @@ +""" +@author: caddiesnew +@file: __init__.py +@time: 2026/8/10 +@description: define_metric resolver。 + +职责: +1. 管理 exact/glob rule 和 concrete state; +2. 在 MetricLogEvent 中解析 concrete definition、执行 step_sync 注入; +3. 为 definition 变化生成 ColumnRecord。 +""" + +from typing import Dict, Optional, Set, Tuple + +from swanlab.proto.swanlab.metric.column.v1.column_pb2 import ( + ColumnClass, + ColumnRecord, + ColumnType, + SectionType, +) +from swanlab.sdk.internal.bus import MetricDefineEvent +from swanlab.sdk.internal.pkg import console, helper + +from .state import ( + ConcreteState, + EffectiveDefinition, + make_effective, +) + +__all__ = [ + "DefinitionResolver", + "EffectiveDefinition", + "ConcreteState", +] + +# 每个 key 的 step 来源记录上限,防止长训练中 _x_step_origins 无限增长。 +# 注入去重只关心"当前 step 是否已注入过",历史 step 的来源记录无需长期保留。 +_MAX_ORIGIN_STEPS_PER_KEY = 16 + +# custom X 值去重的 epsilon:|新值 - 上次值| < 此值视为相同。 +_X_EPSILON = 1e-8 + + +class DefinitionResolver: + """define_metric resolver,生命周期与 Run 一致。 + + 由 BackgroundConsumer 单线程调用,无需加锁。 + """ + + def __init__(self) -> None: + # exact key → 定义快照(define_metric 中不含 '*' 的 key 注册于此) + self._exact_rules: Dict[str, EffectiveDefinition] = {} + # glob pattern(末尾单 '*')→ 定义快照,物化时按最长前缀匹配 + self._glob_rules: Dict[str, EffectiveDefinition] = {} + # (metric_class, key) → 首次 log 时物化的 concrete 定义快照; + # first-writer-wins,之后的 define 不回溯修改已物化条目 + self._concrete: Dict[Tuple[str, str], ConcreteState] = {} + # custom X 源 key → 最近一次真实 log 的 X 值,是跨 step 注入的取值来源 + self._custom_x_cache: Dict[str, float] = {} + # custom X 源 key → {step → 来源标记 REAL / INJECTED / REAL_CONFLICT}; + # 用于同 step 注入去重与"真实 X 晚于注入 X"的冲突判定(记录数有上限,见 _prune_origins) + self._x_step_origins: Dict[str, Dict[int, str]] = {} + # 已发过 X 冲突警告的 key,保证每 key 只警告一次 + self._warned_conflict_x: Set[str] = set() + # 被某 rule 当作 custom X 轴的 key 集合。 + # 只有这些 key 才需要在 _custom_x_cache / _x_step_origins 中注册, + # 避免对未参与 X 关系的海量 key 产生额外开销 + self._custom_x_keys: Set[str] = set() + # Y key → 上次消费的 X 值;连续重复 X 值上的后续 Y 被去重(try_accept_x_value,epsilon 容差) + self._y_last_x_value: Dict[str, float] = {} + + # ── rule 管理 ────────────────────────────────────────────── + + def handle_define(self, event: MetricDefineEvent) -> None: + """处理 define_metric 事件,登记 exact/glob rule。 + + glob 与 exact 都只登记 rule:图表定义遵循 first-writer-wins,由首次 log 时的 + materialize_column 决定;后续 define 不回溯已物化的 key、也不产出 ColumnRecord。 + 同一 key 在 log 前多次 define 仍以最后一次为准(merge/replace 在 rule 上累积)。 + """ + # 1. 按 is_glob 选定 rule 分区 + key = event.key + rules = self._glob_rules if event.is_glob else self._exact_rules + + # 2. 计算 effective:overwrite 从默认值 replace,否则在已有 rule(或默认)上 merge + if event.overwrite: + effective = self._replace_effective(event) + else: + base = rules.get(key) or self._default_effective() + effective = self._merge_effective(base, event) + + # 3. 登记 rule(exact / glob 分区,同 key 后一次 define 覆盖前一次) + rules[key] = effective + + # 4. 记录该 rule 引用的 custom X 源 key,供 update_custom_x_cache 按需登记 + if helper.is_custom_x(effective.x_axis): + self._custom_x_keys.add(effective.x_axis) + + # rule 登记完成。exact 不再立即重新物化已出现的 concrete: + # 图表定义遵循 first-writer-wins,由首次 log 时的 materialize_column 决定。 + # 后续 define 只更新 rule,影响之后首次出现的 key。 + return + + @staticmethod + def _default_effective() -> EffectiveDefinition: + # step_sync 默认值:从未显式指定时等价于 is_custom_x(effective.x_axis) + # (见 _merge_effective / _replace_effective 的默认推导) + return make_effective("_step", None, False, False) + + @staticmethod + def _merge_effective(base: EffectiveDefinition, event: MetricDefineEvent) -> EffectiveDefinition: + x_axis = event.x_axis if event.x_axis is not None else base.x_axis + section_name = event.section_name if event.section_name is not None else base.section_name + hidden = event.hidden if event.hidden is not None else base.hidden + step_sync = event.step_sync + if step_sync is None: + step_sync = base.step_sync + if base.x_axis == "_step" and helper.is_custom_x(x_axis): + # 本次 define 首次为 rule 引入 custom X(x_axis / step_metric)且未显式指定 + # step_sync → 默认 True;显式 False 永远优先(上面分支已处理) + step_sync = True + return make_effective(x_axis, section_name, hidden, step_sync) + + @staticmethod + def _replace_effective(event: MetricDefineEvent) -> EffectiveDefinition: + x_axis = event.x_axis if event.x_axis is not None else "_step" + section_name = event.section_name + hidden = event.hidden if event.hidden is not None else False + # 未指定 step_sync 时重置为默认值:结果 x_axis 为 custom X 则 True,否则 False + step_sync = event.step_sync if event.step_sync is not None else helper.is_custom_x(x_axis) + return make_effective(x_axis, section_name, hidden, step_sync) + + # ── concrete 解析 ────────────────────────────────────────── + + def resolve_concrete(self, key: str, metric_class: str) -> ConcreteState: + """为 log 中的 key 解析 concrete definition。 + + 首次调用时按 exact → glob → automatic 优先级解析并注册 key;之后同一 + (metric_class, key) 的所有调用直接返回首次快照,不再升级或回溯—— + 这同时钉住了 "key 首次 log 后 define 不再生效" 的契约(automatic 状态 + 表示该 key 在无定义下被 log 过,后续 define 无法认领)。 + """ + cache_key = (metric_class, key) + existing = self._concrete.get(cache_key) + if existing is not None: + return existing + + # 以下仅在首次解析时执行 + + # 1. exact 优先 + exact_rule = self._exact_rules.get(key) + if exact_rule: + effective = self._adjust_for_class(exact_rule, metric_class) + state = ConcreteState( + key=key, + metric_class=metric_class, + effective=effective, + source="exact", + ) + self._concrete[cache_key] = state + return state + + # 2. 最长 prefix glob + glob_rule = self._match_glob(key) + if glob_rule: + effective = self._adjust_for_class(glob_rule, metric_class) + state = ConcreteState( + key=key, + metric_class=metric_class, + effective=effective, + source="glob", + ) + self._concrete[cache_key] = state + return state + + # 3. automatic 兜底:无 rule 命中,仅注册钉住状态(materialize 不会为其产出列) + effective = self._adjust_for_class(self._default_effective(), metric_class) + state = ConcreteState( + key=key, + metric_class=metric_class, + effective=effective, + source="automatic", + ) + self._concrete[cache_key] = state + return state + + def _match_glob(self, key: str) -> Optional[EffectiveDefinition]: + """找到最长 prefix glob rule 的定义快照。""" + best: Optional[EffectiveDefinition] = None + best_len = -1 + for pattern, rule in self._glob_rules.items(): + prefix = pattern[:-1] + if key.startswith(prefix) and len(prefix) > best_len: + best = rule + best_len = len(prefix) + return best + + @staticmethod + def _adjust_for_class(effective: EffectiveDefinition, metric_class: str) -> EffectiveDefinition: + """根据 metric class 调整 effective definition。MEDIA 一律回退 _step; + SCALAR 允许自引用(x_axis == key,图像为直线 y=x)。""" + if metric_class == "MEDIA": + return make_effective("_step", effective.section_name, effective.hidden, effective.step_sync) + return effective + + # ── ColumnRecord 生成 ────────────────────────────────────── + + def materialize_column(self, key: str, metric_class: str, column_type: ColumnType) -> Optional[ColumnRecord]: + """物化列定义:每个 (metric_class, key) 仅在首次产出一条 ColumnRecord。 + 当一个 key 首次 log 之后,本 run 内后续的 define 同样不再生效。 + + automatic 来源(无 define 的 key)不产出列——列由 core 收到数据后自动创建, + SDK 侧不重复发送。 + """ + cache_key = (metric_class, key) + state = self._concrete.get(cache_key) + if state is None or state.source == "automatic": + return None + if state.emitted: + return None + state.emitted = True + return self._build_column_record(state, column_type) + + @staticmethod + def _build_column_record(state: ConcreteState, column_type: ColumnType) -> ColumnRecord: + effective = state.effective + col = ColumnRecord( + column_class=ColumnClass.COLUMN_CLASS_CUSTOM, + column_key=state.key, + column_type=column_type, + section_type=SectionType.SECTION_TYPE_PUBLIC, + ) + col.section_name = effective.section_name or "" + if state.metric_class == "SCALAR" and effective.x_axis != "_step": + col.x_axis = effective.x_axis + col.hidden = effective.hidden + return col + + # ── custom X 值去重 ─────────────────────────────────────── + + def try_accept_x_value(self, y_key: str, x_value: float) -> bool: + """X 值去重:与该 Y 上次消费的 X 值相同(epsilon 内)返回 False,否则记录新值并返回 True。 + + 仅抑制连续重复;非单调 X(如 5→6→5)的回退值会被当作新值接受,同一 X 值上可能出现多个 Y 点。 + """ + last = self._y_last_x_value.get(y_key) + if last is not None and abs(last - x_value) < _X_EPSILON: + return False + self._y_last_x_value[y_key] = x_value + return True + + # ── step_sync: custom X cache ───────────────────────────── + + @property + def has_rules(self) -> bool: + """是否登记过任何 exact/glob rule(即 define_metric 是否被调用过)。 + + 为 False 时消费端走纯构建路径:不物化列、不做 custom X 处理, + 列全部由 core 收到数据后自动创建。 + """ + return bool(self._exact_rules or self._glob_rules) + + @property + def has_custom_x(self) -> bool: + """是否登记过任何 custom X 源 key。 + + 为 False 时消费端预扫描(transform + cache 更新)可整体跳过: + update_custom_x_cache 本就以 _custom_x_keys 为门槛,explicit_scalars + 的消费点也全部位于 is_custom_x 分支之后。_custom_x_keys 仅由 + handle_define 添加、clear() 清空,FIFO 队列保证 define 先于后续 log。 + + .. note:: + ``_custom_x_keys`` 只增不减:即使后续 ``overwrite=True`` 把 ``x_axis`` + 重置回 ``_step``,本属性仍返回 True,预扫描对余下训练保持开启。 + 预扫描的 transform 结果会被后续构建复用,额外开销可忽略。 + """ + return bool(self._custom_x_keys) + + def update_custom_x_cache(self, key: str, value: float, step: int) -> None: + """缓存真实 custom X 值并登记 step 来源为 REAL。 + + 仅当 key 是某 rule 的 custom X 源(在 _custom_x_keys 中)时才登记, + 普通指标 key 无需占用 _custom_x_cache / _x_step_origins 内存。 + """ + if helper.is_system_key(key): + return + if key not in self._custom_x_keys: + return + self._custom_x_cache[key] = value + origins = self._x_step_origins.setdefault(key, {}) + if step in origins and origins[step] == "INJECTED": + # Y 先于 X 的场景:此 step 已有注入,真实 X 来了 + # 真实值仍进入 cache,但 Core 去重可能拒绝真实 X 落盘 + if key not in self._warned_conflict_x: + console.warning( + f"Metric '{key}' at step {step}: real value arrives after injected value; " + f"the injected value at this step is kept. Log X before Y in the same step to avoid this." + ) + self._warned_conflict_x.add(key) + origins[step] = "REAL_CONFLICT" + else: + origins[step] = "REAL" + self._prune_origins(origins) + + def get_custom_x(self, key: str) -> Optional[float]: + return self._custom_x_cache.get(key) + + # ── step_sync: X 注入 ────────────────────────────────────── + + def try_inject_x(self, x_key: str, step: int) -> bool: + """尝试为 (x_key, step) 登记注入。已有来源时返回 False。""" + origins = self._x_step_origins.setdefault(x_key, {}) + if step in origins: + return False + origins[step] = "INJECTED" + self._prune_origins(origins) + return True + + @staticmethod + def _prune_origins(origins: Dict[int, str]) -> None: + """限制每个 key 保留的 step 来源记录数量,丢弃最旧的条目。 + + 注入去重只需判断"当前 step 是否已注入过",历史 step 记录无长期价值; + 长训练中不做修剪会导致 _x_step_origins 随 step 数单调增长。 + """ + if len(origins) <= _MAX_ORIGIN_STEPS_PER_KEY: + return + # step 单调递增,按 key(step)升序丢弃最早的条目 + for stale_step in sorted(origins)[: len(origins) - _MAX_ORIGIN_STEPS_PER_KEY]: + del origins[stale_step] + + # ── 清理 ─────────────────────────────────────────────────── + + def clear(self) -> None: + self._exact_rules.clear() + self._glob_rules.clear() + self._concrete.clear() + self._custom_x_cache.clear() + self._x_step_origins.clear() + self._warned_conflict_x.clear() + self._custom_x_keys.clear() + self._y_last_x_value.clear() diff --git a/swanlab/sdk/internal/run/components/consumer/resolver/state.py b/swanlab/sdk/internal/run/components/consumer/resolver/state.py new file mode 100644 index 000000000..1e4a9309f --- /dev/null +++ b/swanlab/sdk/internal/run/components/consumer/resolver/state.py @@ -0,0 +1,48 @@ +""" +@author: caddiesnew +@file: state.py +@time: 2026/8/10 +@description: define_metric resolver 状态数据结构。 +""" + +import sys +from dataclasses import dataclass +from typing import Optional + +__all__ = [ + "EffectiveDefinition", + "ConcreteState", + "make_effective", +] + +# py3.9 兼容:slots 仅 3.10+ 支持 +# 3.9 退化为普通 dict-based dataclass +_SLOTS_KWARGS = {"slots": True} if sys.version_info >= (3, 10) else {} + + +@dataclass(frozen=True, **_SLOTS_KWARGS) +class EffectiveDefinition: + """merge/replace 后的完整 effective 快照。""" + + x_axis: str # X 轴 metric key;"_step" 表示系统默认 + section_name: Optional[str] # 图表分组名称;None 表示使用默认 section + hidden: bool # 是否隐藏该 metric 对应的图表 + step_sync: bool = False # 该 Y 是否在 X 分次 log 时触发 custom X 注入;默认值 = is_custom_x(x_axis),显式指定优先 + + +def make_effective( + x_axis: str, section_name: Optional[str], hidden: bool, step_sync: bool = True +) -> EffectiveDefinition: + """构造一个 EffectiveDefinition。""" + return EffectiveDefinition(x_axis=x_axis, section_name=section_name, hidden=hidden, step_sync=step_sync) + + +@dataclass(**_SLOTS_KWARGS) +class ConcreteState: + """一个已物化的 concrete key 的状态。""" + + key: str # log 中实际出现的 metric key + metric_class: str # resolver identity 的类别:SCALAR 或 MEDIA + effective: EffectiveDefinition # 首次物化或 exact 更新后固定的定义快照 + source: str # 定义来源:exact、glob;automatic 表示无 define 下被 log(仅钉住,不产出列) + emitted: bool = False # 是否已产出过 ColumnRecord(定义不回溯,每 (class, key) 只发一次) diff --git a/swanlab/sdk/internal/run/fmt.py b/swanlab/sdk/internal/run/fmt.py index 7882698e1..66188bb92 100644 --- a/swanlab/sdk/internal/run/fmt.py +++ b/swanlab/sdk/internal/run/fmt.py @@ -10,7 +10,7 @@ from pydantic import ValidationError -from swanlab.sdk.internal.pkg import console, constraints, safe +from swanlab.sdk.internal.pkg import console, constraints, helper, safe from swanlab.sdk.typings.run import FinishType, SaveType from swanlab.sdk.typings.run.column import ScalarXAxisType @@ -163,13 +163,20 @@ def safe_validate_chart_name(name: Optional[str]) -> Optional[str]: def safe_validate_x_axis(x_axis: Optional[ScalarXAxisType]) -> Optional[ScalarXAxisType]: """ - 检查并清洗 x 轴指标名称,如果出现非法字符或长度超过限制,返回 None。 + 校验 ``define_metric`` 的 x 轴值,非法时返回 None。 - :param x_axis: 待检查的 x 轴指标名称 - :return: 清洗后的 x 轴指标名称或 None + - ``None`` 原样返回(未提供,由调用方处理)。 + - ``_step`` / ``_relative_time`` 原样返回(系统 X 轴)。 + - 自定义 key 经 ``MetricKey`` 校验,且拒绝系统指标前缀(``__swanlab__.``)。 + + :param x_axis: 待校验的 x 轴值 + :return: 校验通过的 x 轴值,或 None """ - if x_axis is None: - x_axis = "_step" + # 系统默认与系统 X 轴(_step / _relative_time)原样放行,判定统一走 helper.is_custom_x + if x_axis is None or not helper.is_custom_x(x_axis): + return x_axis + if helper.is_system_key(x_axis): + return None return safe_validate_key(x_axis) diff --git a/swanlab/sdk/typings/core_python/api/upload.py b/swanlab/sdk/typings/core_python/api/upload.py index a4380d62b..d889bf273 100644 --- a/swanlab/sdk/typings/core_python/api/upload.py +++ b/swanlab/sdk/typings/core_python/api/upload.py @@ -50,6 +50,8 @@ "key": Required[str], "type": Required[str], "sectionName": NotRequired[str], + "xAxis": NotRequired[str], + "hidden": NotRequired[bool], }, ) diff --git a/tests/unit/sdk/internal/core_python/test_core_sync.py b/tests/unit/sdk/internal/core_python/test_core_sync.py index 31df15549..5d13f4d7e 100644 --- a/tests/unit/sdk/internal/core_python/test_core_sync.py +++ b/tests/unit/sdk/internal/core_python/test_core_sync.py @@ -199,8 +199,10 @@ def close(self): class FakeMetric: - def __init__(self): + def __init__(self, column_type=None, column=None): self.updated = [] + self.column_type = column_type + self._column = column def ensure_type_match(self, _: ColumnType) -> None: pass @@ -220,12 +222,14 @@ def get(self, key: str) -> Optional[FakeMetric]: return self.metrics.get(key) def define_scalar(self, **kwargs) -> FakeMetric: - metric = FakeMetric() + column = kwargs.get("column") + metric = FakeMetric(column_type=column.column_type if column else None, column=column) self.metrics[kwargs["key"]] = metric return metric def define_media(self, **kwargs) -> FakeMetric: - metric = FakeMetric() + column = kwargs.get("column") + metric = FakeMetric(column_type=column.column_type if column else None, column=column) self.metrics[kwargs["key"]] = metric return metric @@ -581,3 +585,53 @@ def test_confirm_sync_finish_uses_existing_finish_record(tmp_path: Path, monkeyp assert call_kwargs.kwargs["finished_at"] == finished_at assert transport.records == [] assert transport.finished is True + + +# ============================================================ +# define_metric: offline sync parity tests +# ============================================================ + + +def test_sync_replays_column_with_xaxis_hidden(tmp_path: Path): + """带 x_axis/hidden 的 ColumnRecord 在 sync 回放时保留字段。""" + core = CoreSyncPython() + core._ctx = CoreContext(config=make_core_config(tmp_path), mode="sync") + core._ctx.set_online_params("alice", "demo", "project-id", None, "experiment-id") + col = ColumnRecord( + column_key="train/loss", + column_type=ColumnType.COLUMN_TYPE_SCALAR, + x_axis="train/epoch", + section_name="Training", + hidden=True, + ) + core._reader = FakeReader([Record(column=col)]) # type: ignore[assignment] + transport = FakeTransport(core._ctx) + core._transport = transport # type: ignore[assignment] + core._metrics = FakeMetrics() # type: ignore[assignment] + + asyncio.run(core.read()) + + assert len(transport.records) == 1 + replayed = transport.records[0].column + assert replayed.x_axis == "train/epoch" + assert replayed.section_name == "Training" + assert replayed.hidden is True + + +def test_sync_skips_when_key_already_exists(tmp_path: Path): + """两条同 key column → 只上传第一条,第二条跳过(first-writer-wins)。""" + core = CoreSyncPython() + core._ctx = CoreContext(config=make_core_config(tmp_path), mode="sync") + core._ctx.set_online_params("alice", "demo", "project-id", None, "experiment-id") + col1 = ColumnRecord(column_key="loss", column_type=ColumnType.COLUMN_TYPE_SCALAR, section_name="A") + col2 = ColumnRecord(column_key="loss", column_type=ColumnType.COLUMN_TYPE_SCALAR, section_name="B") + core._reader = FakeReader([Record(column=col1), Record(column=col2)]) # type: ignore[assignment] + transport = FakeTransport(core._ctx) + core._transport = transport # type: ignore[assignment] + core._metrics = FakeMetrics() # type: ignore[assignment] + + asyncio.run(core.read()) + + col_records = [r.column for r in transport.records if record_kind(r) == "column"] + assert len(col_records) == 1 + assert col_records[0].section_name == "A" diff --git a/tests/unit/sdk/internal/core_python/transport/test_sender.py b/tests/unit/sdk/internal/core_python/transport/test_sender.py index 165015d3b..74c42d4aa 100644 --- a/tests/unit/sdk/internal/core_python/transport/test_sender.py +++ b/tests/unit/sdk/internal/core_python/transport/test_sender.py @@ -5,12 +5,12 @@ from google.protobuf.timestamp_pb2 import Timestamp from swanlab.exceptions import ApiError -from swanlab.proto.swanlab.metric.column.v1.column_pb2 import ColumnRecord, ColumnType +from swanlab.proto.swanlab.metric.column.v1.column_pb2 import ColumnClass, ColumnRecord, ColumnType from swanlab.proto.swanlab.metric.data.v1.data_pb2 import MediaItem, MediaRecord, MediaValue from swanlab.proto.swanlab.record.v1.record_pb2 import Record from swanlab.proto.swanlab.save.v1.save_pb2 import SaveRecord, SaveType from swanlab.sdk.internal.core_python.context import CoreConfig, CoreContext -from swanlab.sdk.internal.core_python.transport.sender import HttpRecordSender +from swanlab.sdk.internal.core_python.transport.sender import HttpRecordSender, encode_column from swanlab.sdk.internal.core_python.transport.tracker import UploadTracker from swanlab.sdk.internal.core_python.utils import ProgressFileWrapper @@ -575,3 +575,147 @@ def test_resolve_save_source_handles_windows_separators_on_posix(tmp_path: Path) ) ) assert sender._resolve_save_source(custom_rec.save) == custom_fallback + + +# ============================================================ +# encode_column: presence-based encoding +# ============================================================ + + +class TestEncodeColumn: + """encode_column 的 presence-based 编码验证。""" + + def test_inferred_column_only_section_name(self): + """inferred column 只发 sectionName,不发 xAxis/hidden。""" + col = ColumnRecord( + column_key="train/loss", + column_type=ColumnType.COLUMN_TYPE_SCALAR, + section_name="Training", + ) + result = encode_column(col) + assert result == {"key": "train/loss", "type": "FLOAT", "sectionName": "Training"} + + def test_inferred_column_no_section_name(self): + """inferred column 无 sectionName 时只发 key/type。""" + col = ColumnRecord( + column_key="train/loss", + column_type=ColumnType.COLUMN_TYPE_SCALAR, + ) + result = encode_column(col) + assert result == {"key": "train/loss", "type": "FLOAT"} + + def test_column_with_custom_xaxis(self): + """ColumnRecord 带 x_axis → xAxis 发送。""" + col = ColumnRecord( + column_key="train/loss", + column_type=ColumnType.COLUMN_TYPE_SCALAR, + x_axis="train/epoch", + section_name="Training", + ) + result = encode_column(col) + assert result == { + "key": "train/loss", + "type": "FLOAT", + "xAxis": "train/epoch", + "sectionName": "Training", + } + + def test_column_without_xaxis(self): + """ColumnRecord 无 x_axis → xAxis key 不发送。""" + col = ColumnRecord( + column_key="loss", + column_type=ColumnType.COLUMN_TYPE_SCALAR, + section_name="", + ) + result = encode_column(col) + # section_name="" → falsy → 不发送 + assert result == {"key": "loss", "type": "FLOAT"} + + def test_column_with_relative_time_xaxis(self): + """xAxis=_relative_time → 字符串透传,由 Server 转换。""" + col = ColumnRecord( + column_key="loss", + column_type=ColumnType.COLUMN_TYPE_SCALAR, + x_axis="_relative_time", + ) + result = encode_column(col) + assert result == { + "key": "loss", + "type": "FLOAT", + "xAxis": "_relative_time", + } + + def test_column_hidden_true(self): + """hidden=True → hidden key 发送。""" + col = ColumnRecord( + column_key="secret", + column_type=ColumnType.COLUMN_TYPE_SCALAR, + hidden=True, + ) + result = encode_column(col) + assert result == { + "key": "secret", + "type": "FLOAT", + "hidden": True, + } + + def test_column_hidden_false_not_sent(self): + """hidden=False → hidden key 不发送(server 视 absent 为 false)。""" + col = ColumnRecord( + column_key="loss", + column_type=ColumnType.COLUMN_TYPE_SCALAR, + hidden=False, + ) + result = encode_column(col) + assert result == {"key": "loss", "type": "FLOAT"} + + def test_media_column(self): + """media ColumnRecord 无 x_axis(resolver 对 MEDIA 不设 x_axis)。""" + col = ColumnRecord( + column_key="train/img", + column_type=ColumnType.COLUMN_TYPE_IMAGE, + section_name="Images", + ) + result = encode_column(col) + assert result == { + "key": "train/img", + "type": "IMAGE", + "sectionName": "Images", + } + + def test_system_class_overrides_type(self): + """SYSTEM class → type 被覆写为 'SYSTEM'。""" + col = ColumnRecord( + column_key="system/cpu", + column_type=ColumnType.COLUMN_TYPE_SCALAR, + column_class=ColumnClass.COLUMN_CLASS_SYSTEM, + ) + result = encode_column(col) + assert result == {"key": "system/cpu", "type": "SYSTEM"} + + def test_upload_column_end_to_end(self, tmp_path: Path): + """upload_column 端到端:带 xAxis/hidden 的 ColumnRecord 正确编码并发送到 HTTP。""" + sender = _make_sender(tmp_path) + records = [ + Record( + column=ColumnRecord( + column_key="train/loss", + column_type=ColumnType.COLUMN_TYPE_SCALAR, + x_axis="train/epoch", + section_name="Training", + hidden=True, + ) + ) + ] + with patch("swanlab.sdk.internal.core_python.transport.sender.upload_columns") as mock_upload: + sender.upload_column(records) + assert mock_upload.call_count == 1 + series = mock_upload.call_args.kwargs["columns"]["series"] + assert len(series) == 1 + assert series[0] == { + "key": "train/loss", + "type": "FLOAT", + "xAxis": "train/epoch", + "sectionName": "Training", + "hidden": True, + } diff --git a/tests/unit/sdk/internal/run/components/test_builder.py b/tests/unit/sdk/internal/run/components/consumer/builder/test_builder.py similarity index 74% rename from tests/unit/sdk/internal/run/components/test_builder.py rename to tests/unit/sdk/internal/run/components/consumer/builder/test_builder.py index 986a2ccba..10e3a12c8 100644 --- a/tests/unit/sdk/internal/run/components/test_builder.py +++ b/tests/unit/sdk/internal/run/components/consumer/builder/test_builder.py @@ -8,7 +8,8 @@ import pytest from swanlab.proto.swanlab.metric.data.v1.data_pb2 import MediaItem, MediaRecord -from swanlab.sdk.internal.run.components.builder import RecordBuilder +from swanlab.sdk.internal.run.components.consumer.builder import _NON_SCALAR_TYPES, RecordBuilder, is_scalar_value +from swanlab.sdk.internal.run.transforms import Text def _make_media_record(key: str = "test", step: int = 0, items=None) -> MediaRecord: @@ -96,7 +97,36 @@ def test_exact_boundary_size_not_dropped(self, builder): assert result is record def test_exact_boundary_length_not_truncated(self, builder): - """长度恰好等于限制的项不被截断""" + """大小恰好等于限制的项不被截断""" record = _make_media_record(items=[("a.png", 10), ("b.png", 10), ("c.png", 10)]) result = builder._ensure_media_size(record) assert result is record + + +class TestIsScalarValue: + """is_scalar_value 必须与 build_scalar_or_media 的分派结果一致。 + + 消费端用它做预扫描:若新增的媒体注册类型未同步 _NON_SCALAR_TYPES, + 预扫描会把媒体值当标量 transform,导致 media 预处理异常或行为漂移。 + """ + + def test_registry_types_all_covered(self): + """注册表中每个显式类型都必须在 _NON_SCALAR_TYPES 中""" + # Python 3.12 下经 __get__ 包装后的方法不含 registry,需取类上的原始 descriptor + sdm = RecordBuilder.__dict__["build_scalar_or_media"] + registered = {tp for tp in sdm.dispatcher.registry if tp is not object} + assert registered <= set(_NON_SCALAR_TYPES) + + def test_scalars_return_true(self): + assert is_scalar_value(1.5) is True + assert is_scalar_value("str") is True + + def test_media_returns_false(self): + assert is_scalar_value([Text("a"), Text("b")]) is False + assert is_scalar_value(Text("a")) is False + + def test_echarts_returns_false(self): + """默认分支中 echarts 类型会被包装为 ECharts 媒体,非标量""" + from pyecharts.charts import Line + + assert is_scalar_value(Line()) is False diff --git a/tests/unit/sdk/internal/run/components/consumer/resolver/test_resolver.py b/tests/unit/sdk/internal/run/components/consumer/resolver/test_resolver.py new file mode 100644 index 000000000..efa186ab3 --- /dev/null +++ b/tests/unit/sdk/internal/run/components/consumer/resolver/test_resolver.py @@ -0,0 +1,281 @@ +"""DefinitionResolver 单元测试。 + +覆盖:glob 路由(glob 形式合法性由 Run.define_metric 在 API 层校验, +见 test_run_define_metric.py;事件经 ``is_glob`` 显式携带分类结果)、 +metric class 适配(MEDIA 回退 _step、SCALAR 自引用放行)、解析优先级 +(exact > 最长前缀 glob > automatic 钉住)、X 值去重、 +merge/replace 语义(含 step_sync 默认值 = is_custom_x(x_axis))。 +""" + +import pytest + +from swanlab.proto.swanlab.metric.column.v1.column_pb2 import ColumnType +from swanlab.sdk.internal.bus.events import MetricDefineEvent +from swanlab.sdk.internal.run.components.consumer.resolver import DefinitionResolver + + +def _define(resolver: DefinitionResolver, key: str, is_glob: bool = False) -> None: + """对给定 key 触发一次 handle_define(其余字段默认未提供)。""" + resolver.handle_define(MetricDefineEvent(key=key, is_glob=is_glob)) + + +class TestGlobRouting: + """事件 is_glob 的 exact/glob 分区路由。""" + + def test_exact_key_accepted(self): + """无 '*' 的 key 作为 exact rule 注册。""" + r = DefinitionResolver() + _define(r, "train/loss") + assert "train/loss" in r._exact_rules + assert len(r._glob_rules) == 0 + + @pytest.mark.parametrize("pattern", ["train/*", "*"], ids=lambda v: repr(v)) + def test_glob_patterns_route_to_glob_rules(self, pattern): + """末尾单 '*'(含单独的 '*')作为 glob rule 注册。""" + r = DefinitionResolver() + _define(r, pattern, is_glob=True) + assert pattern in r._glob_rules + assert len(r._exact_rules) == 0 + + +class TestSelfReferenceXAxis: + def test_scalar_self_reference_x_axis_is_kept(self): + """SCALAR 允许 x_axis == key(图像为直线 y=x),不再降级为 _step。""" + r = DefinitionResolver() + r.handle_define(MetricDefineEvent("epoch", x_axis="epoch")) + concrete = r.resolve_concrete("epoch", "SCALAR") + assert concrete.effective.x_axis == "epoch" + column = r.materialize_column("epoch", "SCALAR", ColumnType.COLUMN_TYPE_SCALAR) + assert column is not None + assert column.x_axis == "epoch" + + def test_media_x_axis_still_falls_back_to_step(self): + """MEDIA 一律回退 _step(既有行为不变)。""" + r = DefinitionResolver() + r.handle_define(MetricDefineEvent("train/img", x_axis="epoch")) + concrete = r.resolve_concrete("train/img", "MEDIA") + assert concrete.effective.x_axis == "_step" + + +class TestWandbAlignedResolution: + def test_glob_update_does_not_retroactively_change_materialized_concrete(self): + r = DefinitionResolver() + r.handle_define(MetricDefineEvent("train/*", is_glob=True, section_name="A")) + first = r.resolve_concrete("train/loss", "SCALAR") + r.materialize_column("train/loss", "SCALAR", ColumnType.COLUMN_TYPE_SCALAR) + + r.handle_define(MetricDefineEvent("train/*", is_glob=True, section_name="B")) + + assert r.resolve_concrete("train/loss", "SCALAR") is first + assert first.effective.section_name == "A" + assert r.resolve_concrete("train/new_loss", "SCALAR").effective.section_name == "B" + + def test_exact_define_does_not_retroactively_change_glob_concrete(self): + r = DefinitionResolver() + r.handle_define(MetricDefineEvent("train/*", is_glob=True, section_name="Train", hidden=True)) + r.resolve_concrete("train/loss", "SCALAR") + r.materialize_column("train/loss", "SCALAR", ColumnType.COLUMN_TYPE_SCALAR) + + columns = r.handle_define(MetricDefineEvent("train/loss", x_axis="epoch")) + + # 图表定义 first-writer-wins:exact define 不再对已物化 concrete 产出 ColumnRecord + assert columns is None + # rule 仍登记,供之后首次出现的 key 使用 + assert "train/loss" in r._exact_rules + # 已物化的 glob concrete 不被回溯(保持 glob 快照) + effective = r.resolve_concrete("train/loss", "SCALAR").effective + assert effective.x_axis == "_step" + assert effective.section_name == "Train" + assert effective.hidden is True + + def test_multiple_defines_before_log_last_define_wins(self): + """同一 key 在 log 前多次 define,最后一次定义生效(merge 累积在 rule 上)。""" + r = DefinitionResolver() + r.handle_define(MetricDefineEvent("train/loss", x_axis="epoch")) + r.handle_define(MetricDefineEvent("train/loss", x_axis="iter", section_name="Training")) + + concrete = r.resolve_concrete("train/loss", "SCALAR") + assert concrete.effective.x_axis == "iter" + assert concrete.effective.section_name == "Training" + + column = r.materialize_column("train/loss", "SCALAR", ColumnType.COLUMN_TYPE_SCALAR) + assert column is not None + assert column.x_axis == "iter" + assert column.section_name == "Training" + + def test_hidden_three_state_merge_semantics(self): + """hidden 三态:None=保留旧值,True=隐藏,False=显式解除(merge 模式下同样生效)。""" + r = DefinitionResolver() + r.handle_define(MetricDefineEvent("train/loss", hidden=True)) + # 未提供(None)→ 保留 True + r.handle_define(MetricDefineEvent("train/loss", section_name="Train")) + assert r._exact_rules["train/loss"].hidden is True + # 显式 False → merge 下解除隐藏 + r.handle_define(MetricDefineEvent("train/loss", hidden=False)) + assert r._exact_rules["train/loss"].hidden is False + + def test_hidden_replace_resets_to_default(self): + """overwrite=True(replace)下未提供 hidden → 重置为 False。""" + r = DefinitionResolver() + r.handle_define(MetricDefineEvent("train/loss", hidden=True)) + r.handle_define(MetricDefineEvent("train/loss", overwrite=True)) + assert r._exact_rules["train/loss"].hidden is False + + def test_subsequent_log_after_define_does_not_upgrade_automatic_concrete(self): + """log → define → 再 log:concrete 保持首次 automatic 快照,列始终由 core 自动建。""" + r = DefinitionResolver() + # ① 首次 log(无 rule → automatic 钉住,不产出列) + first = r.resolve_concrete("train/loss", "SCALAR") + assert first.source == "automatic" + col1 = r.materialize_column("train/loss", "SCALAR", ColumnType.COLUMN_TYPE_SCALAR) + assert col1 is None # automatic 不产出 ColumnRecord,列由 core 收到数据后自动创建 + + # ② define 登记 exact rule(x_axis=epoch),不产出 ColumnRecord + columns = r.handle_define(MetricDefineEvent("train/loss", x_axis="epoch")) + assert columns is None + assert "train/loss" in r._exact_rules # rule 仍登记 + + # ③ 再 log:resolve_concrete 返回冻结的首次 concrete,不升级为 exact + second = r.resolve_concrete("train/loss", "SCALAR") + assert second is first # 同一对象 + assert second.source == "automatic" # 仍是 automatic,未被 exact 认领 + assert second.effective.x_axis == "_step" # x_axis 未变 + + # ④ materialize_column 仍不产出(automatic 来源永不物化) + col2 = r.materialize_column("train/loss", "SCALAR", ColumnType.COLUMN_TYPE_SCALAR) + assert col2 is None + + +class TestXValueDedup: + """try_accept_x_value:custom X 值 first-writer-wins 去重。""" + + def test_new_value_accepted(self): + """首次出现的值与其后的不同值均被接受。""" + r = DefinitionResolver() + assert r.try_accept_x_value("acc", 5.0) is True + assert r.try_accept_x_value("acc", 6.0) is True + + def test_same_value_rejected(self): + r = DefinitionResolver() + r.try_accept_x_value("acc", 5.0) + assert r.try_accept_x_value("acc", 5.0) is False + + def test_non_monotonic_x_is_not_deduplicated(self): + """去重仅对比上一个 X 值,假设 X 单调递增。 + + 若 X 非单调(5→6→5),回退值会被当作新值接受,同一 X 值可能出现多个 Y。 + 这不是设计保证——``try_accept_x_value`` 仅抑制**连续重复**的 X 值; + "每个 X 值仅保留首个 Y" 的完整保证依赖调用方保证 custom X 单调递增。 + """ + r = DefinitionResolver() + r.try_accept_x_value("acc", 5.0) + r.try_accept_x_value("acc", 6.0) + # 单调假设下不会出现 5→6→5;此处仅记录连续去重的实现行为 + assert r.try_accept_x_value("acc", 5.0) is True + + @pytest.mark.parametrize( + "delta,expected", + [(1e-9, False), (1e-7, True)], + ids=["within-epsilon", "outside-epsilon"], + ) + def test_epsilon_threshold(self, delta, expected): + """epsilon(1e-8)内的微小差异视为相同,超出视为不同。""" + r = DefinitionResolver() + r.try_accept_x_value("loss", 1.0) + assert r.try_accept_x_value("loss", 1.0 + delta) is expected + + def test_x_value_zero_accepted_then_rejected(self): + """X=0.0 不受 falsy 影响。""" + r = DefinitionResolver() + assert r.try_accept_x_value("acc", 0.0) is True + assert r.try_accept_x_value("acc", 0.0) is False + + def test_per_y_key_independent(self): + """不同 Y key 各自独立去重。""" + r = DefinitionResolver() + r.try_accept_x_value("loss", 5.0) + assert r.try_accept_x_value("acc", 5.0) is True # acc 与 loss 独立 + + +class TestStepSyncMergeReplace: + """step_sync 在 define 时的默认值、merge、replace。 + + 默认值语义:从未显式指定时,step_sync = is_custom_x(effective.x_axis), + 即提供 x_axis / step_metric 时默认 True,否则 False;显式 False 永远优先。 + """ + + def test_default_step_sync_is_true(self): + r = DefinitionResolver() + r.handle_define(MetricDefineEvent("loss", x_axis="epoch")) + assert r.resolve_concrete("loss", "SCALAR").effective.step_sync is True + + def test_default_step_sync_false_without_x_axis(self): + r = DefinitionResolver() + r.handle_define(MetricDefineEvent("loss", section_name="Train")) + assert r.resolve_concrete("loss", "SCALAR").effective.step_sync is False + + def test_attaching_x_axis_later_defaults_step_sync_true(self): + """先无 x_axis 定义、后补 x_axis 且未显式指定 step_sync → 默认 True。""" + r = DefinitionResolver() + r.handle_define(MetricDefineEvent("loss", section_name="Train")) + r.handle_define(MetricDefineEvent("loss", x_axis="epoch")) + effective = r.resolve_concrete("loss", "SCALAR").effective + assert effective.x_axis == "epoch" + assert effective.step_sync is True + + def test_explicit_false_is_kept(self): + r = DefinitionResolver() + r.handle_define(MetricDefineEvent("loss", x_axis="epoch", step_sync=False)) + assert r.resolve_concrete("loss", "SCALAR").effective.step_sync is False + + def test_merge_updates_step_sync_without_clearing_x_axis(self): + r = DefinitionResolver() + r.handle_define(MetricDefineEvent("loss", x_axis="epoch")) + r.handle_define(MetricDefineEvent("loss", step_sync=False)) + effective = r.resolve_concrete("loss", "SCALAR").effective + assert effective.x_axis == "epoch" + assert effective.step_sync is False + + def test_merge_omitted_step_sync_keeps_previous(self): + r = DefinitionResolver() + r.handle_define(MetricDefineEvent("loss", x_axis="epoch", step_sync=False)) + r.handle_define(MetricDefineEvent("loss", section_name="Train")) + effective = r.resolve_concrete("loss", "SCALAR").effective + assert effective.step_sync is False + assert effective.section_name == "Train" + + def test_overwrite_unspecified_resets_step_sync_to_true(self): + r = DefinitionResolver() + r.handle_define(MetricDefineEvent("loss", x_axis="epoch", step_sync=False)) + r.handle_define(MetricDefineEvent("loss", x_axis="epoch", overwrite=True)) + assert r.resolve_concrete("loss", "SCALAR").effective.step_sync is True + + def test_overwrite_unspecified_resets_step_sync_to_false_without_x_axis(self): + """overwrite 未指定 step_sync 时重置为默认值:无 custom X 则 False。""" + r = DefinitionResolver() + r.handle_define(MetricDefineEvent("loss", x_axis="epoch")) + r.handle_define(MetricDefineEvent("loss", overwrite=True)) + effective = r.resolve_concrete("loss", "SCALAR").effective + assert effective.x_axis == "_step" + assert effective.step_sync is False + + def test_overwrite_explicit_false_is_kept(self): + r = DefinitionResolver() + r.handle_define(MetricDefineEvent("loss", x_axis="epoch")) + r.handle_define(MetricDefineEvent("loss", x_axis="epoch", step_sync=False, overwrite=True)) + assert r.resolve_concrete("loss", "SCALAR").effective.step_sync is False + + def test_glob_step_sync_false_applies_to_matching_key(self): + r = DefinitionResolver() + r.handle_define(MetricDefineEvent("train/*", is_glob=True, x_axis="epoch", step_sync=False)) + assert r.resolve_concrete("train/loss", "SCALAR").effective.step_sync is False + + def test_define_after_materialize_does_not_change_step_sync(self): + r = DefinitionResolver() + r.handle_define(MetricDefineEvent("loss", x_axis="epoch")) + first = r.resolve_concrete("loss", "SCALAR") + r.materialize_column("loss", "SCALAR", ColumnType.COLUMN_TYPE_SCALAR) + r.handle_define(MetricDefineEvent("loss", step_sync=False)) + second = r.resolve_concrete("loss", "SCALAR") + assert second is first + assert second.effective.step_sync is True diff --git a/tests/unit/sdk/internal/run/components/consumer/test_consumer_log.py b/tests/unit/sdk/internal/run/components/consumer/test_consumer_log.py new file mode 100644 index 000000000..0e9f07198 --- /dev/null +++ b/tests/unit/sdk/internal/run/components/consumer/test_consumer_log.py @@ -0,0 +1,95 @@ +"""BackgroundConsumer 指标日志分流测试:无 define 走 _log_plain,有 define 走 _log_define。 + +关键契约: +- 无 define:纯构建 data record,不产出 ColumnRecord(列由 core 收到数据后自动创建); +- 有 define:define 过的 key 首次 log 物化一次列,后续 log 不再产出; +- key 首次 log 后到达的 define 无法认领该 key(automatic 钉住)。 +""" + +import queue +from pathlib import Path +from types import SimpleNamespace +from typing import cast +from unittest.mock import MagicMock + +from google.protobuf.timestamp_pb2 import Timestamp + +from swanlab.sdk.internal.bus.events import MetricDefineEvent, MetricLogEvent +from swanlab.sdk.internal.context import RunContext +from swanlab.sdk.internal.pkg import console as pkg_console +from swanlab.sdk.internal.run.components.consumer import BackgroundConsumer + + +def _make_consumer(tmp_path: Path) -> BackgroundConsumer: + ctx = SimpleNamespace( + core=MagicMock(), + callbacker=MagicMock(), + media_dir=tmp_path, + config=SimpleNamespace(settings=SimpleNamespace(core=SimpleNamespace(section_rule=0))), + ) + run_ctx = cast(RunContext, cast(object, ctx)) + return BackgroundConsumer(run_ctx, queue.Queue()) + + +def _log_event(data: dict, step: int = 1) -> MetricLogEvent: + return MetricLogEvent(data=data, step=step, timestamp=Timestamp(seconds=step)) + + +class TestPlainPath: + """无 define 路径:纯构建,零列产出。""" + + def test_no_define_no_columns(self, tmp_path: Path): + consumer = _make_consumer(tmp_path) + consumer._handle_event(_log_event({"loss": 0.5})) + + assert len(consumer._scalar_batch) == 1 + assert consumer._scalar_batch[0].key == "loss" + assert consumer._column_batch == [] + + def test_forged_system_key_filtered_with_warning(self, tmp_path: Path, monkeypatch): + warning = MagicMock() + monkeypatch.setattr(pkg_console, "warning", warning) + consumer = _make_consumer(tmp_path) + consumer._handle_event(_log_event({"__swanlab__.cpu": 1, "loss": 2})) + + warning.assert_called_once() + assert len(consumer._scalar_batch) == 1 + assert consumer._scalar_batch[0].key == "loss" + assert consumer._column_batch == [] + + +class TestDefinePath: + """有 define 路径:define 过的 key 物化一次列。""" + + def test_define_materializes_column_once(self, tmp_path: Path): + consumer = _make_consumer(tmp_path) + consumer._handle_event(MetricDefineEvent(key="train/loss", section_name="Train")) + + consumer._handle_event(_log_event({"train/loss": 0.5})) + consumer._handle_event(_log_event({"train/loss": 0.4}, step=2)) + + assert len(consumer._column_batch) == 1 + assert consumer._column_batch[0].column_key == "train/loss" + assert consumer._column_batch[0].section_name == "Train" + assert len(consumer._scalar_batch) == 2 + + def test_undefined_key_in_define_run_stays_columnless(self, tmp_path: Path): + """同一 run 内未 define 的 key 仍不产出列(core 自动建)。""" + consumer = _make_consumer(tmp_path) + consumer._handle_event(MetricDefineEvent(key="train/loss", section_name="Train")) + + consumer._handle_event(_log_event({"train/loss": 0.5, "other": 1})) + + assert len(consumer._column_batch) == 1 + assert consumer._column_batch[0].column_key == "train/loss" + assert len(consumer._scalar_batch) == 2 + + def test_late_define_cannot_claim_logged_key(self, tmp_path: Path): + """key 首次 log 后到达的 define 无法认领该 key(automatic 钉住)。""" + consumer = _make_consumer(tmp_path) + consumer._handle_event(_log_event({"loss": 0.5})) # 无 define,先 log + consumer._handle_event(MetricDefineEvent(key="loss", section_name="Train")) + consumer._handle_event(_log_event({"loss": 0.4}, step=2)) + + assert consumer._column_batch == [] + assert len(consumer._scalar_batch) == 2 diff --git a/tests/unit/sdk/internal/run/components/test_consumer_save.py b/tests/unit/sdk/internal/run/components/consumer/test_consumer_save.py similarity index 95% rename from tests/unit/sdk/internal/run/components/test_consumer_save.py rename to tests/unit/sdk/internal/run/components/consumer/test_consumer_save.py index 580ea307a..3ccc2999e 100644 --- a/tests/unit/sdk/internal/run/components/test_consumer_save.py +++ b/tests/unit/sdk/internal/run/components/consumer/test_consumer_save.py @@ -27,8 +27,11 @@ def _make_consumer(tmp_path: Path, batch_size: int = 100): source_path=str(tmp_path / "model.pt"), policy=SavePolicy.SAVE_POLICY_NOW, ) + consumer = BackgroundConsumer(cast(RunContext, cast(object, ctx)), queue.Queue(), batch_size=batch_size) + # builder 由 consumer 内部创建,测试以 mock 替换以拦截 build_save 结果 + consumer._builder = builder return ( - BackgroundConsumer(cast(RunContext, cast(object, ctx)), queue.Queue(), builder, batch_size=batch_size), + consumer, core, builder, ) diff --git a/tests/unit/sdk/internal/run/components/consumer/test_consumer_step_sync.py b/tests/unit/sdk/internal/run/components/consumer/test_consumer_step_sync.py new file mode 100644 index 000000000..aa0a92a7e --- /dev/null +++ b/tests/unit/sdk/internal/run/components/consumer/test_consumer_step_sync.py @@ -0,0 +1,436 @@ +"""BackgroundConsumer step_sync 注入路径测试。""" + +import math +import queue +from pathlib import Path +from types import SimpleNamespace +from typing import cast +from unittest.mock import MagicMock + +import pytest +from google.protobuf.timestamp_pb2 import Timestamp + +from swanlab.proto.swanlab.metric.column.v1.column_pb2 import ColumnType +from swanlab.sdk.internal.bus.events import MetricDefineEvent, MetricLogEvent +from swanlab.sdk.internal.context import RunContext +from swanlab.sdk.internal.run.components.consumer import BackgroundConsumer +from swanlab.sdk.internal.run.transforms import Text + + +def _make_consumer(tmp_path: Path) -> BackgroundConsumer: + ctx = SimpleNamespace( + core=MagicMock(), + callbacker=MagicMock(), + media_dir=tmp_path, + config=SimpleNamespace(settings=SimpleNamespace(core=SimpleNamespace(section_rule=0))), + ) + run_ctx = cast(RunContext, cast(object, ctx)) + return BackgroundConsumer(run_ctx, queue.Queue()) + + +def _timestamp(seconds: int) -> Timestamp: + return Timestamp(seconds=seconds) + + +def _log(consumer: BackgroundConsumer, data: dict, step: int, seconds: int) -> None: + consumer._handle_event(MetricLogEvent(data=data, step=step, timestamp=_timestamp(seconds))) + + +def _records_by_key(consumer: BackgroundConsumer, start: int = 0): + records = consumer._scalar_batch[start:] + return {key: [record for record in records if record.key == key] for key in {record.key for record in records}} + + +def test_step_sync_injects_cached_x_with_current_y_step_and_timestamp(tmp_path: Path): + consumer = _make_consumer(tmp_path) + consumer._handle_event(MetricDefineEvent("train/loss", x_axis="epoch")) + _log(consumer, {"epoch": 3}, step=0, seconds=10) + start = len(consumer._scalar_batch) + + _log(consumer, {"train/loss": 0.8}, step=1, seconds=20) + + records = _records_by_key(consumer, start) + assert set(records) == {"epoch", "train/loss"} + assert len(records["epoch"]) == 1 + assert records["epoch"][0].value.number == 3 + assert records["epoch"][0].step == 1 + assert records["epoch"][0].timestamp == _timestamp(20) + assert records["train/loss"][0].value.number == pytest.approx(0.8) + assert records["train/loss"][0].step == 1 + assert records["train/loss"][0].timestamp == _timestamp(20) + # 只有 define 过的 train/loss 产出列;epoch 未 define,列由 core 收到数据后自动创建 + assert [column.column_key for column in consumer._column_batch] == ["train/loss"] + + +@pytest.mark.parametrize( + "data", + [ + {"epoch": 3, "train/loss": 0.8}, + {"train/loss": 0.8, "epoch": 3}, + ], + ids=["x-before-y", "y-before-x"], +) +def test_step_sync_does_not_inject_when_event_contains_x_regardless_of_order(tmp_path: Path, data: dict): + consumer = _make_consumer(tmp_path) + consumer._handle_event(MetricDefineEvent("train/loss", x_axis="epoch")) + + _log(consumer, data, step=1, seconds=20) + + records = _records_by_key(consumer) + assert set(records) == {"epoch", "train/loss"} + assert len(records["epoch"]) == 1 + assert records["epoch"][0].value.number == 3 + + +def test_step_sync_injects_shared_x_only_once_for_multiple_y_in_same_event(tmp_path: Path): + consumer = _make_consumer(tmp_path) + consumer._handle_event(MetricDefineEvent("train/acc", x_axis="epoch")) + consumer._handle_event(MetricDefineEvent("train/loss", x_axis="epoch")) + _log(consumer, {"epoch": 3}, step=0, seconds=10) + start = len(consumer._scalar_batch) + + _log(consumer, {"train/acc": 0.9, "train/loss": 0.1}, step=1, seconds=20) + + records = _records_by_key(consumer, start) + assert set(records) == {"epoch", "train/acc", "train/loss"} + assert len(records["epoch"]) == 1 + assert records["epoch"][0].value.number == 3 + + +def test_step_sync_does_not_inject_shared_x_twice_at_same_step(tmp_path: Path): + consumer = _make_consumer(tmp_path) + consumer._handle_event(MetricDefineEvent("train/acc", x_axis="epoch")) + consumer._handle_event(MetricDefineEvent("train/loss", x_axis="epoch")) + _log(consumer, {"epoch": 3}, step=0, seconds=10) + start = len(consumer._scalar_batch) + + _log(consumer, {"train/acc": 0.9}, step=1, seconds=20) + _log(consumer, {"train/loss": 0.1}, step=1, seconds=21) + + records = _records_by_key(consumer, start) + assert len(records["epoch"]) == 1 + assert records["epoch"][0].timestamp == _timestamp(20) + + +def test_real_x_after_injected_x_warns_once_and_updates_next_step_cache(tmp_path: Path, monkeypatch): + warnings = [] + monkeypatch.setattr( + "swanlab.sdk.internal.run.components.consumer.resolver.console.warning", + lambda message: warnings.append(message), + ) + consumer = _make_consumer(tmp_path) + consumer._handle_event(MetricDefineEvent("train/loss", x_axis="epoch")) + _log(consumer, {"epoch": 3}, step=0, seconds=10) + _log(consumer, {"train/loss": 0.8}, step=1, seconds=20) + + _log(consumer, {"epoch": 4}, step=1, seconds=21) + _log(consumer, {"epoch": 4}, step=1, seconds=22) + start = len(consumer._scalar_batch) + _log(consumer, {"train/loss": 0.7}, step=2, seconds=30) + + records = _records_by_key(consumer, start) + assert len(warnings) == 1 + assert "real value arrives after injected value" in warnings[0] + assert len(records["epoch"]) == 1 + assert records["epoch"][0].value.number == 4 + assert records["epoch"][0].step == 2 + assert records["epoch"][0].timestamp == _timestamp(30) + + +# ============================================================ +# custom X first-writer-wins:同一 X 值上首次 Y 值为准 +# ============================================================ + + +def test_duplicate_x_drops_second_y(tmp_path: Path): + """X 只 log 一次,两次 Y log → 第二次 Y 丢弃(注入相同 X 值)。""" + consumer = _make_consumer(tmp_path) + consumer._handle_event(MetricDefineEvent("acc", x_axis="x1")) + _log(consumer, {"x1": 5}, step=0, seconds=10) + start = len(consumer._scalar_batch) + + _log(consumer, {"acc": 0.9}, step=1, seconds=20) # inject x1=5.0 → accepted + _log(consumer, {"acc": 1.0}, step=2, seconds=30) # inject x1=5.0 → rejected + + acc_records = [r for r in consumer._scalar_batch[start:] if r.key == "acc"] + assert len(acc_records) == 1 + assert acc_records[0].value.number == pytest.approx(0.9) + + +def test_different_x_values_accepted(tmp_path: Path): + """X 更新后,新 X 值的 Y 被接受。""" + consumer = _make_consumer(tmp_path) + consumer._handle_event(MetricDefineEvent("acc", x_axis="x1")) + _log(consumer, {"x1": 5}, step=0, seconds=10) + _log(consumer, {"acc": 0.9}, step=1, seconds=20) # inject x1=5.0 → accepted + _log(consumer, {"x1": 6}, step=2, seconds=30) + start = len(consumer._scalar_batch) + + _log(consumer, {"acc": 1.0}, step=3, seconds=40) # inject x1=6.0 → accepted + + acc_records = [r for r in consumer._scalar_batch[start:] if r.key == "acc"] + assert len(acc_records) == 1 + assert acc_records[0].value.number == pytest.approx(1.0) + + +def test_explicit_same_x_value_rejected(tmp_path: Path): + """显式 log X=5 在不同 step,Y 仍按 X 值去重。""" + consumer = _make_consumer(tmp_path) + consumer._handle_event(MetricDefineEvent("acc", x_axis="x1")) + + _log(consumer, {"x1": 5, "acc": 0.9}, step=0, seconds=10) # accepted + _log(consumer, {"x1": 5, "acc": 1.0}, step=1, seconds=20) # rejected (same X value) + + acc_records = [r for r in consumer._scalar_batch if r.key == "acc"] + assert len(acc_records) == 1 + assert acc_records[0].value.number == pytest.approx(0.9) + + +def test_multiple_y_sharing_x_independent(tmp_path: Path): + """loss 和 acc 共享 x1,各自独立去重。""" + consumer = _make_consumer(tmp_path) + consumer._handle_event(MetricDefineEvent("loss", x_axis="x1")) + consumer._handle_event(MetricDefineEvent("acc", x_axis="x1")) + _log(consumer, {"x1": 5}, step=0, seconds=10) + _log(consumer, {"loss": 0.5, "acc": 0.9}, step=1, seconds=20) # both accepted + start = len(consumer._scalar_batch) + + _log(consumer, {"loss": 0.4, "acc": 0.8}, step=2, seconds=30) # both rejected + + new_records = consumer._scalar_batch[start:] + loss_records = [r for r in new_records if r.key == "loss"] + acc_records = [r for r in new_records if r.key == "acc"] + assert len(loss_records) == 0 + assert len(acc_records) == 0 + + +def test_no_custom_x_no_dedup(tmp_path: Path): + """无 custom X 轴时,不触发去重。""" + consumer = _make_consumer(tmp_path) + + _log(consumer, {"acc": 0.9}, step=0, seconds=10) + _log(consumer, {"acc": 1.0}, step=1, seconds=20) + + acc_records = [r for r in consumer._scalar_batch if r.key == "acc"] + assert len(acc_records) == 2 + + +# ============================================================ +# M1:非有限 X 值不应被注入的旧值覆盖 +# 注入跳过条件必须看 event.data(用户是否提供过该 key),而非 explicit_scalars +# (后者已被 isfinite 过滤);否则 {**data, **injected} 会用旧值覆盖用户的 nan +# ============================================================ + + +def test_nan_x_value_not_overwritten_by_injection(tmp_path: Path): + """X 为 NaN 时,落盘的应是 nan(用户显式值),而非被注入的旧值 3.0。""" + consumer = _make_consumer(tmp_path) + consumer._handle_event(MetricDefineEvent("loss", x_axis="epoch")) + _log(consumer, {"epoch": 3}, step=0, seconds=10) + _log(consumer, {"epoch": float("nan"), "loss": 0.8}, step=1, seconds=20) + + epoch_records = [r for r in consumer._scalar_batch if r.key == "epoch"] + # 第二条 epoch 必须是用户显式的 nan,而非注入的 3.0 + assert math.isnan(epoch_records[-1].value.number) + # 不应出现 step=1 处 value=3 的注入孤儿点 + injected_orphan = [r for r in epoch_records if r.step == 1 and r.value.number == 3.0] + assert injected_orphan == [] + + +# ============================================================ +# S2:Y 被 X 去重丢弃时不留孤儿注入点、不误标 INJECTED +# 注入 commit(try_inject_x + record)必须推迟到确认至少一个存活 Y 之后 +# ============================================================ + + +def test_dropped_y_leaves_no_orphan_injected_x(tmp_path: Path): + """Y 被 X 值去重丢弃时,不为其注入孤儿 X 点(注入只对存活 Y commit)。""" + consumer = _make_consumer(tmp_path) + consumer._handle_event(MetricDefineEvent("loss", x_axis="epoch")) + _log(consumer, {"epoch": 5}, step=0, seconds=10) + _log(consumer, {"loss": 0.9}, step=1, seconds=20) # inject epoch=5@1,Y 存活 + start = len(consumer._scalar_batch) + + _log(consumer, {"loss": 0.8}, step=2, seconds=30) # epoch=5 连续重复 → Y 丢弃 + + new_records = consumer._scalar_batch[start:] + # loss 被去重丢弃 + assert [r for r in new_records if r.key == "loss"] == [] + # 未为被丢的 Y 注入孤儿 epoch 点(step=2 处无 epoch) + assert [r for r in new_records if r.key == "epoch" and r.step == 2] == [] + + +def test_dropped_y_no_spurious_real_conflict_warning(tmp_path: Path, monkeypatch): + """Y 被丢弃后未标 INJECTED,真实 X 晚到同 step 不触发误 REAL_CONFLICT。""" + warnings = [] + monkeypatch.setattr( + "swanlab.sdk.internal.run.components.consumer.resolver.console.warning", + lambda message: warnings.append(message), + ) + consumer = _make_consumer(tmp_path) + consumer._handle_event(MetricDefineEvent("loss", x_axis="epoch")) + _log(consumer, {"epoch": 5}, step=0, seconds=10) + _log(consumer, {"loss": 0.9}, step=1, seconds=20) # inject epoch=5@1(Y 存活,标 INJECTED@1) + _log(consumer, {"loss": 0.8}, step=2, seconds=30) # epoch=5 重复 → Y 丢弃,不标 INJECTED@2 + _log(consumer, {"epoch": 6}, step=2, seconds=31) # 真实 X@2,未标 INJECTED@2 → 无误告警 + + conflict = [w for w in warnings if "real value arrives after injected value" in w] + assert conflict == [] + + +def test_step_sync_false_does_not_inject_cached_x(tmp_path: Path): + """step_sync=False:X/Y 分次 log 时不注入 X。""" + consumer = _make_consumer(tmp_path) + consumer._handle_event(MetricDefineEvent("train/loss", x_axis="epoch", step_sync=False)) + _log(consumer, {"epoch": 3}, step=0, seconds=10) + start = len(consumer._scalar_batch) + + _log(consumer, {"train/loss": 0.8}, step=1, seconds=20) + + records = _records_by_key(consumer, start) + assert set(records) == {"train/loss"} + assert records["train/loss"][0].value.number == pytest.approx(0.8) + assert records["train/loss"][0].step == 1 + + +def test_step_sync_false_does_not_drop_y_from_cached_x(tmp_path: Path): + """False 不用跨 step cache 去重:event 不含 X 时 Y 保留。""" + consumer = _make_consumer(tmp_path) + consumer._handle_event(MetricDefineEvent("acc", x_axis="x1", step_sync=False)) + _log(consumer, {"x1": 5, "acc": 0.9}, step=0, seconds=10) + start = len(consumer._scalar_batch) + + _log(consumer, {"acc": 1.0}, step=1, seconds=20) + + acc_records = [r for r in consumer._scalar_batch[start:] if r.key == "acc"] + assert len(acc_records) == 1 + assert acc_records[0].value.number == pytest.approx(1.0) + assert [r for r in consumer._scalar_batch[start:] if r.key == "x1"] == [] + + +def test_step_sync_false_still_drops_y_when_event_has_duplicate_x(tmp_path: Path): + """False 仍按本 event 显式 X 去重。""" + consumer = _make_consumer(tmp_path) + consumer._handle_event(MetricDefineEvent("acc", x_axis="x1", step_sync=False)) + _log(consumer, {"x1": 5, "acc": 0.9}, step=0, seconds=10) + _log(consumer, {"x1": 5, "acc": 1.0}, step=1, seconds=20) + + acc_records = [r for r in consumer._scalar_batch if r.key == "acc"] + assert len(acc_records) == 1 + assert acc_records[0].value.number == pytest.approx(0.9) + + +def test_step_sync_false_same_step_real_x_does_not_inject_or_drop(tmp_path: Path): + """False 不用本 step 真实 X:分次 log 时 Y 保留、不注入。""" + consumer = _make_consumer(tmp_path) + consumer._handle_event(MetricDefineEvent("acc", x_axis="x1", step_sync=False)) + _log(consumer, {"x1": 5}, step=1, seconds=10) + start = len(consumer._scalar_batch) + + _log(consumer, {"acc": 0.9}, step=1, seconds=20) + + records = _records_by_key(consumer, start) + assert set(records) == {"acc"} + assert records["acc"][0].value.number == pytest.approx(0.9) + + +def test_step_sync_false_sibling_does_not_block_true_sibling_inject(tmp_path: Path): + """False 的 Y 不触发注入;同 event 的 True 兄弟仍注入共享 X。""" + consumer = _make_consumer(tmp_path) + consumer._handle_event(MetricDefineEvent("train/acc", x_axis="epoch", step_sync=False)) + consumer._handle_event(MetricDefineEvent("train/loss", x_axis="epoch")) + _log(consumer, {"epoch": 3}, step=0, seconds=10) + start = len(consumer._scalar_batch) + + _log(consumer, {"train/acc": 0.9, "train/loss": 0.1}, step=1, seconds=20) + + records = _records_by_key(consumer, start) + assert set(records) == {"epoch", "train/acc", "train/loss"} + assert len(records["epoch"]) == 1 + assert records["epoch"][0].value.number == 3 + assert records["epoch"][0].step == 1 + + +def test_step_sync_false_alone_does_not_inject_even_with_true_sibling_defined(tmp_path: Path): + """仅 False 的 Y 出现时不注入,即使存在 step_sync=True 的兄弟 rule。""" + consumer = _make_consumer(tmp_path) + consumer._handle_event(MetricDefineEvent("train/acc", x_axis="epoch", step_sync=False)) + consumer._handle_event(MetricDefineEvent("train/loss", x_axis="epoch")) + _log(consumer, {"epoch": 3}, step=0, seconds=10) + start = len(consumer._scalar_batch) + + _log(consumer, {"train/acc": 0.9}, step=1, seconds=20) + + records = _records_by_key(consumer, start) + assert set(records) == {"train/acc"} + + +def test_overwrite_resets_step_sync_to_true_before_first_log(tmp_path: Path): + """overwrite=True 且未指定 step_sync 时恢复默认注入。""" + consumer = _make_consumer(tmp_path) + consumer._handle_event(MetricDefineEvent("train/loss", x_axis="epoch", step_sync=False)) + consumer._handle_event(MetricDefineEvent("train/loss", x_axis="epoch", overwrite=True)) + _log(consumer, {"epoch": 3}, step=0, seconds=10) + start = len(consumer._scalar_batch) + + _log(consumer, {"train/loss": 0.8}, step=1, seconds=20) + + records = _records_by_key(consumer, start) + assert "epoch" in records + assert records["epoch"][0].value.number == 3 + assert records["epoch"][0].step == 1 + + +def test_self_reference_x_axis_keeps_first_log_wins(tmp_path: Path): + """自引用 x_axis(x_axis == key):正常落盘自身值,且 first-writer-wins 语义不变。 + + - 自引用 key 恒随 Y 出现,不触发 step_sync 注入; + - 连续重复 X 值按既有 epsilon 规则丢弃该 Y 点; + - 二次 define 不改变已冻结 concrete(仍自引用)。 + """ + consumer = _make_consumer(tmp_path) + consumer._handle_event(MetricDefineEvent("epoch", x_axis="epoch")) + + _log(consumer, {"epoch": 1}, step=0, seconds=10) + _log(consumer, {"epoch": 2}, step=1, seconds=20) + _log(consumer, {"epoch": 2}, step=2, seconds=30) # 连续重复 X → 丢弃 + _log(consumer, {"epoch": 3}, step=3, seconds=40) + + records = _records_by_key(consumer) + assert set(records) == {"epoch"} + assert [r.value.number for r in records["epoch"]] == [1, 2, 3] + assert [r.step for r in records["epoch"]] == [0, 1, 3] + + # ColumnRecord 仅一条,x_axis 为自身 + columns = [c for c in consumer._column_batch if c.column_key == "epoch"] + assert len(columns) == 1 + assert columns[0].x_axis == "epoch" + + # 二次 define(first-writer-wins:concrete 已冻结,不再产出 ColumnRecord) + consumer._handle_event(MetricDefineEvent("epoch", x_axis="_step")) + concrete = consumer._resolver.resolve_concrete("epoch", "SCALAR") + assert concrete.effective.x_axis == "epoch" + assert [c for c in consumer._column_batch if c.column_key == "epoch"] == [columns[0]] + + +# ============================================================ +# 快路径:无 custom X 规则时预扫描整体跳过, +# scalar/media 混合 log 仍正常走 builder 路径(media 落盘、列定义产出) +# ============================================================ + + +def test_no_custom_x_mixed_media_and_scalar_log(tmp_path: Path): + consumer = _make_consumer(tmp_path) + assert consumer._resolver.has_custom_x is False + + _log(consumer, {"train/loss": 0.8, "train/text": Text("hello")}, step=0, seconds=10) + + scalars = [r for r in consumer._scalar_batch if r.key == "train/loss"] + assert len(scalars) == 1 + assert scalars[0].value.number == pytest.approx(0.8) + media = [r for r in consumer._media_batch if r.key == "train/text"] + assert len(media) == 1 + assert media[0].type == ColumnType.COLUMN_TYPE_TEXT + assert media[0].value.items[0].filename.endswith(".txt") + # 无 define:SDK 不产出 ColumnRecord,列由 core 收到数据后自动创建 + assert consumer._column_batch == [] diff --git a/tests/unit/sdk/internal/run/test_run_define_metric.py b/tests/unit/sdk/internal/run/test_run_define_metric.py new file mode 100644 index 000000000..e3ab97d52 --- /dev/null +++ b/tests/unit/sdk/internal/run/test_run_define_metric.py @@ -0,0 +1,100 @@ +"""Run.define_metric 输入校验分支的最小集成测试。 + +key 的清洗/截断/强转语义由 test_run_fmt.py 的 TestValidateKey 覆盖, +此处固化接线行为:清洗后的 key 进入事件、非法 key / 非法 glob 报错且不发出事件、 +glob 分类结果(is_glob)随事件显式下发。 +""" + +import threading +from unittest.mock import MagicMock + +import pytest + +from swanlab.sdk.internal.bus.events import MetricDefineEvent +from swanlab.sdk.internal.pkg import fork +from swanlab.sdk.internal.run import Run, fmt + + +class _MockRun: + """最小化的 Run 替身,供非绑定方法测试使用""" + + def __init__(self): + self._api_lock = threading.RLock() + self._init_pid = fork.current_pid() + self.alive = True + self._components = MagicMock() + + +@pytest.fixture(autouse=True) +def _reset_warned_keys(): + """fmt._WARNED_KEYS 是进程级缓存,用例前后重置避免跨用例污染。""" + fmt._WARNED_KEYS.clear() + yield + fmt._WARNED_KEYS.clear() + + +def _emitted_key(mock_run: _MockRun) -> str: + """断言事件已发出并返回其 key。""" + mock_run._components.emitter.emit.assert_called_once() + event = mock_run._components.emitter.emit.call_args[0][0] + assert isinstance(event, MetricDefineEvent) + return event.key + + +class TestDefineMetricKeyValidation: + def test_sanitized_key_emitted(self, monkeypatch): + """非法边缘字符被清洗后再进入事件,规则与 log 侧的规范 key 对齐。""" + monkeypatch.setattr(fmt.console, "warning", MagicMock()) + mock_run = _MockRun() + Run.define_metric(mock_run, "train/loss/") # type: ignore + assert _emitted_key(mock_run) == "train/loss" + + def test_invalid_key_rejected_without_emit(self, monkeypatch): + """清洗后为空的 key 报错且不发出任何事件。""" + error = MagicMock() + monkeypatch.setattr(fmt.console, "error", error) + mock_run = _MockRun() + Run.define_metric(mock_run, "///") # type: ignore + error.assert_called_once() + mock_run._components.emitter.emit.assert_not_called() + + +class TestDefineMetricGlobValidation: + """glob 模式校验与 is_glob 分类(校验在 API 层同步完成,resolver 不再解析字符串)。""" + + @pytest.mark.parametrize( + "key,is_glob", + [("train/*", True), ("*", True), ("train/loss", False)], + ids=["trailing-star", "bare-star", "exact"], + ) + def test_glob_classification_emitted(self, key, is_glob, monkeypatch): + """合法 key 原样进入事件,is_glob 按末尾单 '*' 分类。""" + monkeypatch.setattr(fmt.console, "error", MagicMock()) + mock_run = _MockRun() + Run.define_metric(mock_run, key) # type: ignore + mock_run._components.emitter.emit.assert_called_once() + event = mock_run._components.emitter.emit.call_args[0][0] + assert isinstance(event, MetricDefineEvent) + assert event.key == key + assert event.is_glob is is_glob + + @pytest.mark.parametrize( + "bad", + [ + "*loss", # '*' 在开头 + "train/*/loss", # '*' 在中间 + "train/**", # 末尾两个 '*' + "**", # 多个 '*' + "a*b", # '*' 在中间且非末尾 + "a***b", # 多个 '*' + ], + ids=lambda v: repr(v), + ) + def test_invalid_glob_rejected_without_emit(self, bad, monkeypatch): + """非法 glob 报错且不发出任何事件。""" + error = MagicMock() + monkeypatch.setattr(fmt.console, "error", error) + mock_run = _MockRun() + Run.define_metric(mock_run, bad) # type: ignore + error.assert_called_once() + mock_run._components.emitter.emit.assert_not_called() From 42e173bed906b45869d56ca7993a8db3e1c6e2c5 Mon Sep 17 00:00:00 2001 From: Kang Li <79990647+SAKURA-CAT@users.noreply.github.com> Date: Sat, 29 Aug 2026 20:27:13 +0800 Subject: [PATCH 05/11] chore: deprecate unused fmt validators (#1763) --- swanlab/deprecated/fmt.py | 52 +++++++++++++++++++ .../components/consumer/resolver/__init__.py | 9 ++-- swanlab/sdk/internal/run/fmt.py | 32 +----------- 3 files changed, 57 insertions(+), 36 deletions(-) create mode 100644 swanlab/deprecated/fmt.py diff --git a/swanlab/deprecated/fmt.py b/swanlab/deprecated/fmt.py new file mode 100644 index 000000000..56c6ec6fb --- /dev/null +++ b/swanlab/deprecated/fmt.py @@ -0,0 +1,52 @@ +""" +@author: caddiesnew +@description: run.fmt 中随 define_scalar 移除而失去调用方的校验函数,v0.11 删除 +""" + +import warnings + +from typing_extensions import deprecated + +from swanlab.sdk.internal.pkg import constraints + + +@deprecated("`safe_validate_chart_name()` has no callers and will be removed in v0.11.") +def safe_validate_chart_name(name): + """ + 检查并清洗图表名称,如果出现非法字符或长度超过限制,返回 None。 + + :param name: 待检查的图表名称 + :return: 清洗后的图表名称或 None + """ + warnings.warn( + "`safe_validate_chart_name()` has no callers and will be removed in v0.11.", + FutureWarning, + stacklevel=2, + ) + if name is None: + return None + try: + return constraints.ta_chart_name.validate_python(name) + except Exception: + return None + + +@deprecated("`safe_validate_color()` has no callers and will be removed in v0.11.") +def safe_validate_color(color): + """ + 检查并清洗颜色字符串格式,必须是#开头的十六进制颜色代码 + + :param color: 待检查的颜色字符串 + :return: 清洗后的颜色字符串或 None + """ + warnings.warn( + "`safe_validate_color()` has no callers and will be removed in v0.11.", + FutureWarning, + stacklevel=2, + ) + if color is None: + return None + try: + return constraints.ta_hex_color.validate_python(color) + except Exception: + return None diff --git a/swanlab/sdk/internal/run/components/consumer/resolver/__init__.py b/swanlab/sdk/internal/run/components/consumer/resolver/__init__.py index 1351b4f9e..63621f529 100644 --- a/swanlab/sdk/internal/run/components/consumer/resolver/__init__.py +++ b/swanlab/sdk/internal/run/components/consumer/resolver/__init__.py @@ -133,12 +133,11 @@ def _replace_effective(event: MetricDefineEvent) -> EffectiveDefinition: # ── concrete 解析 ────────────────────────────────────────── def resolve_concrete(self, key: str, metric_class: str) -> ConcreteState: - """为 log 中的 key 解析 concrete definition。 + """为 log 中的 key 解析 concrete 定义,首次调用时注册并冻结快照。 - 首次调用时按 exact → glob → automatic 优先级解析并注册 key;之后同一 - (metric_class, key) 的所有调用直接返回首次快照,不再升级或回溯—— - 这同时钉住了 "key 首次 log 后 define 不再生效" 的契约(automatic 状态 - 表示该 key 在无定义下被 log 过,后续 define 无法认领)。 + 解析优先级:exact > 最长前缀 glob > automatic。首次调用(log)按命中结果注册 ConcreteState; + 之后同一 (metric_class, key) 的调用直接返回该快照,不再受后续 define_metric 影响。 + 这意味着一旦log,后续 define_metric 对已出现的 key 不再回溯修改图表定义,遵循 first-writer-wins的原则。 """ cache_key = (metric_class, key) existing = self._concrete.get(cache_key) diff --git a/swanlab/sdk/internal/run/fmt.py b/swanlab/sdk/internal/run/fmt.py index 66188bb92..5ace85bd1 100644 --- a/swanlab/sdk/internal/run/fmt.py +++ b/swanlab/sdk/internal/run/fmt.py @@ -146,21 +146,6 @@ def safe_validate_name(name: Optional[str]) -> Optional[str]: return None -def safe_validate_chart_name(name: Optional[str]) -> Optional[str]: - """ - 检查并清洗图表名称,如果出现非法字符或长度超过限制,返回 None。 - - :param name: 待检查的图表名称 - :return: 清洗后的图表名称或 None - """ - if name is None: - return None - try: - return constraints.ta_chart_name.validate_python(name) - except ValidationError: - return None - - def safe_validate_x_axis(x_axis: Optional[ScalarXAxisType]) -> Optional[ScalarXAxisType]: """ 校验 ``define_metric`` 的 x 轴值,非法时返回 None。 @@ -177,22 +162,7 @@ def safe_validate_x_axis(x_axis: Optional[ScalarXAxisType]) -> Optional[ScalarXA return x_axis if helper.is_system_key(x_axis): return None - return safe_validate_key(x_axis) - - -def safe_validate_color(color: Optional[str]) -> Optional[str]: - """ - 检查并清洗颜色字符串格式,必须是#开头的十六进制颜色代码 - - :param color: 待检查的颜色字符串 - :return: 清洗后的颜色字符串或 None - """ - if color is None: - return None - try: - return constraints.ta_hex_color.validate_python(color) - except ValidationError: - return None + return safe_validate_key(x_axis) def safe_validate_state(state: FinishType) -> Optional[FinishType]: From 44cd523ddcb1010ed86e9aeb82a11362935b41f5 Mon Sep 17 00:00:00 2001 From: CaddiesNew <50736785+Nexisato@users.noreply.github.com> Date: Tue, 1 Sep 2026 12:00:14 +0800 Subject: [PATCH 06/11] fix: all custom keys return none (#1764) --- README.md | 2 +- README_EN.md | 2 +- README_JP.md | 2 +- README_RU.md | 2 +- .../run/components/consumer/__init__.py | 5 ++-- swanlab/sdk/internal/run/fmt.py | 2 +- .../consumer/test_consumer_step_sync.py | 30 +++++++++++++++++++ 7 files changed, 38 insertions(+), 7 deletions(-) diff --git a/README.md b/README.md index 98f7d7b12..11c32b146 100644 --- a/README.md +++ b/README.md @@ -320,7 +320,7 @@ import swanlab # 初始化一个新的swanlab实验 swanlab.init( project="my-first-ml", - config={'learning-rate': 0.003}, + config={"learning-rate": 0.003}, ) # 记录指标 diff --git a/README_EN.md b/README_EN.md index a80e0eb47..2321ff389 100644 --- a/README_EN.md +++ b/README_EN.md @@ -318,7 +318,7 @@ import swanlab # Initialize a new SwanLab experiment swanlab.init( project="my-first-ml", - config={'learning-rate': 0.003}, + config={"learning-rate": 0.003}, ) # Log metrics diff --git a/README_JP.md b/README_JP.md index b01e0a433..7539f9b35 100644 --- a/README_JP.md +++ b/README_JP.md @@ -298,7 +298,7 @@ import swanlab # 新しいSwanLab実験を初期化 swanlab.init( project="my-first-ml", - config={'learning-rate': 0.003}, + config={"learning-rate": 0.003}, ) # 指標を記録 diff --git a/README_RU.md b/README_RU.md index 6d5c4d75f..a0a0a53f9 100644 --- a/README_RU.md +++ b/README_RU.md @@ -302,7 +302,7 @@ import swanlab # Инициализация нового эксперимента SwanLab swanlab.init( project="my-first-ml", - config={'learning-rate': 0.003}, + config={"learning-rate": 0.003}, ) # Запись метрик diff --git a/swanlab/sdk/internal/run/components/consumer/__init__.py b/swanlab/sdk/internal/run/components/consumer/__init__.py index 94cbc6ef9..fd780ff00 100644 --- a/swanlab/sdk/internal/run/components/consumer/__init__.py +++ b/swanlab/sdk/internal/run/components/consumer/__init__.py @@ -266,10 +266,11 @@ def _log_define(self, event: MetricLogEvent, data: Dict[str, Any]) -> None: # 3.3 如果用户在本次 event 里显式 log 了 X 值,则不注入 if x_axis in data or x_axis in candidate_x: continue - # 3.4 从 cache 取最近真实 X 值 + # 3.4 从 cache 取最近真实 X 值;X 尚无任何真实值时注入默认 0, + # 使 X 序列从首个绑定 Y 的 step 起出现,而非等到 X 自身首次 log cached_x = self._resolver.get_custom_x(x_axis) if cached_x is None: - continue + cached_x = 0.0 candidate_x[x_axis] = cached_x # 4. 把本次 log 的每个 key 变成 record,包括: diff --git a/swanlab/sdk/internal/run/fmt.py b/swanlab/sdk/internal/run/fmt.py index 5ace85bd1..125a40d2a 100644 --- a/swanlab/sdk/internal/run/fmt.py +++ b/swanlab/sdk/internal/run/fmt.py @@ -162,7 +162,7 @@ def safe_validate_x_axis(x_axis: Optional[ScalarXAxisType]) -> Optional[ScalarXA return x_axis if helper.is_system_key(x_axis): return None - return safe_validate_key(x_axis) + return safe_validate_key(x_axis) def safe_validate_state(state: FinishType) -> Optional[FinishType]: diff --git a/tests/unit/sdk/internal/run/components/consumer/test_consumer_step_sync.py b/tests/unit/sdk/internal/run/components/consumer/test_consumer_step_sync.py index aa0a92a7e..8137b06db 100644 --- a/tests/unit/sdk/internal/run/components/consumer/test_consumer_step_sync.py +++ b/tests/unit/sdk/internal/run/components/consumer/test_consumer_step_sync.py @@ -112,6 +112,36 @@ def test_step_sync_does_not_inject_shared_x_twice_at_same_step(tmp_path: Path): assert records["epoch"][0].timestamp == _timestamp(20) +def test_step_sync_injects_default_zero_before_first_real_x(tmp_path: Path): + """X 从未 log 过时,绑定 Y 触发注入默认 0,X 序列从首个绑定 Y 的 step 出现(对齐 wandb)。""" + consumer = _make_consumer(tmp_path) + consumer._handle_event(MetricDefineEvent("train/acc", x_axis="epoch")) + start = len(consumer._scalar_batch) + + _log(consumer, {"train/acc": 0.9}, step=0, seconds=10) + + records = _records_by_key(consumer, start) + assert set(records) == {"train/acc", "epoch"} + assert records["epoch"][0].value.number == 0 + assert records["epoch"][0].step == 0 + assert records["epoch"][0].timestamp == _timestamp(10) + + +def test_step_sync_default_zero_at_each_y_step_then_forward_fill_after_real_x(tmp_path: Path): + """无值窗口内每个 Y event 在各自 step 注入 0;真实 X 到达后恢复取最近真实值。""" + consumer = _make_consumer(tmp_path) + consumer._handle_event(MetricDefineEvent("train/acc", x_axis="epoch")) + consumer._handle_event(MetricDefineEvent("train/loss", x_axis="epoch")) + + _log(consumer, {"train/acc": 0.9}, step=0, seconds=10) + _log(consumer, {"train/loss": 0.1}, step=1, seconds=20) + _log(consumer, {"epoch": 3}, step=2, seconds=30) + _log(consumer, {"train/acc": 0.8}, step=3, seconds=40) + + epoch_records = [r for r in consumer._scalar_batch if r.key == "epoch"] + assert [(r.step, r.value.number) for r in epoch_records] == [(0, 0), (1, 0), (2, 3), (3, 3)] + + def test_real_x_after_injected_x_warns_once_and_updates_next_step_cache(tmp_path: Path, monkeypatch): warnings = [] monkeypatch.setattr( From a1c3c76305d78c7671b937917e7efb9f96e36f06 Mon Sep 17 00:00:00 2001 From: Kang Li <79990647+SAKURA-CAT@users.noreply.github.com> Date: Tue, 1 Sep 2026 15:02:47 +0800 Subject: [PATCH 07/11] chore: sync __all__ exports and drop dead sync cleanup (#1765) --- swanlab/sdk/internal/core_python/pkg/builder/__init__.py | 1 + swanlab/sdk/internal/core_python/sync.py | 1 - swanlab/sdk/internal/pkg/adapter/__init__.py | 2 +- 3 files changed, 2 insertions(+), 2 deletions(-) diff --git a/swanlab/sdk/internal/core_python/pkg/builder/__init__.py b/swanlab/sdk/internal/core_python/pkg/builder/__init__.py index 6d68c42db..2b19aba60 100644 --- a/swanlab/sdk/internal/core_python/pkg/builder/__init__.py +++ b/swanlab/sdk/internal/core_python/pkg/builder/__init__.py @@ -29,6 +29,7 @@ "build_auto_column", "build_resume_column", "build_save_record", + "build_column_record", ] diff --git a/swanlab/sdk/internal/core_python/sync.py b/swanlab/sdk/internal/core_python/sync.py index 9be3d559f..4802f2646 100644 --- a/swanlab/sdk/internal/core_python/sync.py +++ b/swanlab/sdk/internal/core_python/sync.py @@ -323,7 +323,6 @@ def confirm_sync_finish(self) -> ConfirmSyncFinishResponse: state=finish_record.state, finished_at=finish_record.finished_at, ) - self._pending_online_finish_record = None return ConfirmSyncFinishResponse(success=True, message="OK") # 如果仅仅是与后端同步出现问题,则换一个让用户安心一些的提示信息 return ConfirmSyncFinishResponse( diff --git a/swanlab/sdk/internal/pkg/adapter/__init__.py b/swanlab/sdk/internal/pkg/adapter/__init__.py index a5360e7a7..77385f230 100644 --- a/swanlab/sdk/internal/pkg/adapter/__init__.py +++ b/swanlab/sdk/internal/pkg/adapter/__init__.py @@ -15,7 +15,7 @@ from .bimap import BiMap -__all__ = ["resume", "medium", "state", "level", "policy", "memory_unit", "accelerator_vendor"] +__all__ = ["resume", "medium", "column", "state", "level", "policy", "memory_unit", "accelerator_vendor"] resume = BiMap( From 133fe56d06a2e7f8c6c149393ac66fca1fde015b Mon Sep 17 00:00:00 2001 From: CaddiesNew <50736785+Nexisato@users.noreply.github.com> Date: Sun, 6 Sep 2026 18:20:06 +0800 Subject: [PATCH 08/11] feat: add custom x_axis param (#1768) --- swanlab/api/__init__.py | 2 +- swanlab/api/column.py | 4 +- swanlab/api/experiment.py | 65 ++-- swanlab/api/helper/__init__.py | 91 +++++ swanlab/api/helper/extractor.py | 299 +++++++++++++++ swanlab/api/helper/request.py | 152 ++++++++ swanlab/api/{ => helper}/utils.py | 29 ++ swanlab/api/metric.py | 580 ++++++++++-------------------- swanlab/api/project.py | 2 +- swanlab/api/self_hosted.py | 2 +- swanlab/api/series.py | 10 +- swanlab/api/typings/__init__.py | 6 + swanlab/api/typings/common.py | 28 +- swanlab/api/typings/metric.py | 26 +- swanlab/api/user.py | 2 +- swanlab/api/workspace.py | 2 +- swanlab/cli/api/experiment.py | 31 +- tests/unit/api/test_api.py | 112 +++++- tests/unit/api/test_extractor.py | 167 +++++++++ tests/unit/api/test_utils.py | 114 +++++- 20 files changed, 1283 insertions(+), 441 deletions(-) create mode 100644 swanlab/api/helper/__init__.py create mode 100644 swanlab/api/helper/extractor.py create mode 100644 swanlab/api/helper/request.py rename swanlab/api/{ => helper}/utils.py (87%) create mode 100644 tests/unit/api/test_extractor.py diff --git a/swanlab/api/__init__.py b/swanlab/api/__init__.py index 6327240d7..b3c6f33f4 100644 --- a/swanlab/api/__init__.py +++ b/swanlab/api/__init__.py @@ -18,6 +18,7 @@ from .base import ApiClientContext, BaseEntity from .column import Column, Columns from .experiment import Experiment, Experiments +from .helper import validate_api_path, validate_non_empty_string from .project import Project, Projects from .self_hosted import SelfHosted from .series import Series @@ -30,7 +31,6 @@ PaginatedQuery, ) from .user import User -from .utils import validate_api_path, validate_non_empty_string from .workspace import Workspace, Workspaces diff --git a/swanlab/api/column.py b/swanlab/api/column.py index f49a83ecc..1c24241cc 100644 --- a/swanlab/api/column.py +++ b/swanlab/api/column.py @@ -13,6 +13,7 @@ from typing import Any, Callable, Dict, Iterator, Optional, cast from swanlab.api.base import ApiClientContext, BaseEntity +from swanlab.api.helper import get_properties, parse_column_data_type, resolve_run_path, validate_column_params from swanlab.api.typings.column import ApiColumnType from swanlab.api.typings.common import ( ApiColumnClassLiteral, @@ -21,7 +22,6 @@ ApiResponseType, PaginatedQuery, ) -from swanlab.api.utils import get_properties, parse_column_data_type, resolve_run_path, validate_column_params class Column(BaseEntity): @@ -183,6 +183,8 @@ def metric( root_pro_id=self._root_pro_id, root_exp_id=self._root_exp_id, created_at=self._exp_created_at, + # Column 已废弃:不参与 x 轴扩展,显式冻结 step 语义 + x_axis="step", ) return metric.json() diff --git a/swanlab/api/experiment.py b/swanlab/api/experiment.py index bc89eaceb..5d8f0c175 100644 --- a/swanlab/api/experiment.py +++ b/swanlab/api/experiment.py @@ -12,7 +12,16 @@ from typing_extensions import deprecated from swanlab.api.base import ApiClientContext, BaseEntity -from swanlab.api.typings import ApiResponseType +from swanlab.api.helper import ( + get_properties, + parse_timestamp_ms, + resolve_run_path, + validate_filter, + validate_group, + validate_sort, + validate_update_active, +) +from swanlab.api.typings import ApiMetricXAxisParam, ApiResponseType from swanlab.api.typings.common import ( ApiColumnClassLiteral, ApiColumnDataTypeLiteral, @@ -28,15 +37,6 @@ ApiExperimentType, ) from swanlab.api.typings.user import ApiUserType -from swanlab.api.utils import ( - get_properties, - parse_timestamp_ms, - resolve_run_path, - validate_filter, - validate_group, - validate_sort, - validate_update_active, -) from swanlab.sdk.internal.pkg import console from swanlab.utils.time import parse_timestamp_s @@ -243,18 +243,24 @@ def metrics( ignore_timestamp: bool = False, all: bool = False, range_query: Optional[Union[Dict[str, Any], RangeQuery]] = None, + x_axis: ApiMetricXAxisParam = "step", ) -> Dict[str, Any]: """ Fetch scalar metrics (e.g. loss, acc) with three query modes: 1. **Sampled** (default) — server-side LTTB downsampling, up to ``sample`` data points 2. **Full** — ``all=True``, no sampling limit - 3. **Range** — filter by step / timestamp / recent time window via ``range_query`` + 3. **Range** — filter by step / timestamp / custom x value domain via ``range_query`` .. note:: Modes 2 and 3 (``all`` / ``range_query``) download full-resolution CSV data and perform range filtering client-side. Each metric point contains ``step``, ``value``, - and ``timestamp`` (if available). + and ``timestamp`` (if available); under a custom x axis it additionally carries + ``index`` (the custom x value). Under a custom x axis without timestamp-based filtering + (``last`` or ``type="timestamp"``), missing cells are emitted as ``NaN`` placeholders + instead of being skipped, so the per-key lists align by position and can be zipped directly. + Note that the CLI (orjson) renders ``NaN`` / infinities as ``null``; the sampled mode + (default) drops such points server-side instead. :param keys: Metric keys to fetch, e.g. ``["loss", "acc"]`` :param sample: Max sampled data points (default 1500, max 1500). Ignored when ``all`` or ``range_query`` is set. @@ -262,6 +268,8 @@ def metrics( :param all: If True, fetch full-resolution data without sampling limit :param range_query: Range filter — accepts a ``RangeQuery`` object or a plain dict. Only supported for SCALAR metrics. + :param x_axis: X axis of the ``index`` values — ``"step"`` (default), ``"time"`` / + ``"relative_time"`` (built-in axes), or any other non-empty string as a custom x column key. --- @@ -270,9 +278,13 @@ def metrics( ========== ============ ================================================= Field Type Description ========== ============ ================================================= - ``type`` ``str`` Filter axis: ``"step"`` (default) or ``"timestamp"`` - ``start`` ``int`` Lower bound (inclusive); None = from beginning - ``end`` ``int`` Upper bound (inclusive); None = to end + ``type`` ``str`` Filter axis: ``"step"`` (default), ``"timestamp"``, or + ``"custom"`` (custom x value domain; requires a custom x axis) + ``start`` ``float`` Lower bound (inclusive); None = from beginning. Stored + as float (int input accepted). Must be a non-negative + integer value for step/timestamp; any finite float (incl. + negative) for custom + ``end`` ``float`` Upper bound (inclusive); None = to end. Same rules as start ``last`` ``int`` Last N milliseconds (mutually exclusive with start/end) ``head`` ``int`` First N data points (mutually exclusive with tail) ``tail`` ``int`` Last N data points (mutually exclusive with head) @@ -288,19 +300,27 @@ def metrics( **Examples — progressive** - 1. Default sampled query:: + 1. Default sampled query (step x axis):: exp.metrics(keys=["loss", "acc"]) - 2. Filter by step range:: + 2. Explicit custom x axis (e.g. plot loss against epoch):: + + exp.metrics(keys=["loss"], x_axis="epoch") + + 3. Filter by step range:: exp.metrics(keys=["loss"], range_query={"start": 100, "end": 500}) - 3. Step range + first 50 points:: + 4. Filter by custom x value domain (floats / negatives allowed):: + + exp.metrics(keys=["loss"], x_axis="lr", range_query={"type": "custom", "start": 1e-4, "end": 1e-3}) + + 5. Step range + first 50 points:: exp.metrics(keys=["loss"], range_query={"start": 0, "end": 500, "head": 50}) - 4. Filter by timestamp (Unix ms, auto-padded if < 13 digits):: + 6. Filter by timestamp (Unix ms, auto-padded if < 13 digits):: exp.metrics(keys=["loss"], range_query={ "type": "timestamp", @@ -308,15 +328,15 @@ def metrics( "end": 1715773200000, }) - 5. Last 5 minutes:: + 7. Last 5 minutes:: exp.metrics(keys=["loss"], range_query={"last": 300_000}) - 6. Last 5 minutes + first 20 points:: + 8. Last 5 minutes + first 20 points:: exp.metrics(keys=["loss"], range_query={"last": 300_000, "head": 20}) - 7. Last 30 data points:: + 9. Last 30 data points:: exp.metrics(keys=["loss"], range_query={"tail": 30}) """ @@ -347,6 +367,7 @@ def metrics( root_exp_id=self.root_exp_id, created_at=self.created_at_ts, experiment_name=self.name, + x_axis=x_axis, ).json() def summary( diff --git a/swanlab/api/helper/__init__.py b/swanlab/api/helper/__init__.py new file mode 100644 index 000000000..0524bad73 --- /dev/null +++ b/swanlab/api/helper/__init__.py @@ -0,0 +1,91 @@ +""" +@author: caddiesnew +@file: __init__.py +@time: 2026/9/3 +@description: SwanLab OpenAPI 实体层公共辅助函数集合(统一出口,详见各领域模块) +""" + +from swanlab.api.helper.extractor import ( + RELATIVE_TIME_AXIS, + SCALAR_STATISTIC_FIELDS, + STEP_AXIS, + TIME_AXIS, + align_entries_by_key, + axis_request_params, + builtin_x_axis, + extract_first, + extract_value_stats, + merge_value_stats, + stream_export_csv, +) +from swanlab.api.helper.request import ( + build_column_ref, + build_export_payload, + build_media_items, + build_media_payload, + build_scalar_payload, + fetch_file_presigned_urls, + fetch_presigned_urls, +) +from swanlab.api.helper.utils import ( + get_properties, + parse_column_data_type, + parse_timestamp_ms, + resolve_run_path, + strip_dict, + validate_api_path, + validate_column_params, + validate_filter, + validate_group, + validate_metric_keys, + validate_metric_log_level, + validate_metric_type, + validate_non_empty_string, + validate_project_name, + validate_sort, + validate_update_active, + validate_visibility, + validate_x_axis, +) + +__all__ = [ + # extractor —— 标量数据提取与解析(轴解析、统计合并、CSV 流式解析) + "SCALAR_STATISTIC_FIELDS", + "RELATIVE_TIME_AXIS", + "STEP_AXIS", + "TIME_AXIS", + "align_entries_by_key", + "axis_request_params", + "builtin_x_axis", + "extract_first", + "extract_value_stats", + "merge_value_stats", + "stream_export_csv", + # request —— House 请求体构建与预签名资源获取 + "build_column_ref", + "build_export_payload", + "build_media_items", + "build_media_payload", + "build_scalar_payload", + "fetch_file_presigned_urls", + "fetch_presigned_urls", + # utils —— 参数校验与通用工具 + "get_properties", + "parse_column_data_type", + "parse_timestamp_ms", + "resolve_run_path", + "strip_dict", + "validate_api_path", + "validate_column_params", + "validate_filter", + "validate_group", + "validate_metric_keys", + "validate_metric_log_level", + "validate_metric_type", + "validate_non_empty_string", + "validate_project_name", + "validate_sort", + "validate_update_active", + "validate_visibility", + "validate_x_axis", +] diff --git a/swanlab/api/helper/extractor.py b/swanlab/api/helper/extractor.py new file mode 100644 index 000000000..519525cb5 --- /dev/null +++ b/swanlab/api/helper/extractor.py @@ -0,0 +1,299 @@ +""" +@author: caddiesnew +@file: extractor.py +@time: 2026/9/3 +@description: 标量指标数据提取辅助 — X 轴解析/分组、统计值合并、导出 CSV 流式解析 +""" + +import math +from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple + +from swanlab.api.typings.common import ApiResponseType, RangeQuery +from swanlab.api.typings.metric import ApiMetricXAxisType +from swanlab.sdk.internal.pkg import console, safe + +if TYPE_CHECKING: + from swanlab.sdk.internal.pkg.client import Client + +# --------------------------------------------------------------------------- +# X 轴解析与分组 +# --------------------------------------------------------------------------- + +# 回显形式:描述 metrics[].index 的实际语义 +STEP_AXIS: ApiMetricXAxisType = {"type": "step"} +TIME_AXIS: ApiMetricXAxisType = {"type": "SYSTEM", "key": "time"} +RELATIVE_TIME_AXIS: ApiMetricXAxisType = {"type": "SYSTEM", "key": "relative_time"} + + +def builtin_x_axis(x_axis: str) -> Optional[ApiMetricXAxisType]: + """显式内置轴字面量 → 回显形式;自定义 key 字符串返回 None。""" + return { + "step": STEP_AXIS, + "time": TIME_AXIS, + "relative_time": RELATIVE_TIME_AXIS, + }.get(x_axis) + + +def axis_request_params(axis: ApiMetricXAxisType) -> Tuple[str, str]: + """解析后的轴 → House 折线请求参数 ``(xType, xKey)``。 + + CUSTOM 轴走 ``xKey``(服务端 xKey != "" 时自动切换 range/join 语义,xType 不参与); + SYSTEM 内置轴走 ``xType``;step 轴即现状默认。 + """ + axis_type = axis.get("type", "step") + if axis_type == "CUSTOM": + return ("step", axis.get("key", "")) + key = axis.get("key", "") + if axis_type == "SYSTEM" and key in ("time", "relative_time"): + return (key, "") + return ("step", "") + + +# --------------------------------------------------------------------------- +# 响应对齐与统计值合并 +# --------------------------------------------------------------------------- + +SCALAR_STATISTIC_FIELDS = ("min", "max", "avg", "median", "latest") + + +def extract_first(resp: ApiResponseType) -> Optional[Dict[str, Any]]: + """从列表型 API 响应中提取第一个元素,失败返回 None。""" + if resp.ok and isinstance(resp.data, list) and resp.data: + return resp.data[0] + return None + + +def align_entries_by_key(entries: List[Dict[str, Any]]) -> Dict[str, Dict[str, Any]]: + """将后端返回的列表按 ``key`` 字段映射为 dict,应对后端可能省略列或乱序返回。""" + indexed: Dict[str, Dict[str, Any]] = {} + for entry in entries: + if not isinstance(entry, dict): + continue + key = entry.get("key") + if isinstance(key, str) and key: + indexed[key] = entry + return indexed + + +def merge_value_stats( + step_list: List[Dict[str, Any]], + time_list: List[Dict[str, Any]], + keys: List[str], +) -> List[Dict[str, Any]]: + """合并 step/time 两种 x_type 的 value stat 响应为 per-key 统计字典。""" + step_by_key = align_entries_by_key(step_list) + time_by_key = align_entries_by_key(time_list) + merged: List[Dict[str, Any]] = [] + for key in keys: + step_entry = step_by_key.get(key, {}) + time_entry = time_by_key.get(key, {}) + entry: Dict[str, Any] = {} + for field in SCALAR_STATISTIC_FIELDS: + step_val = step_entry.get(field) + time_val = time_entry.get(field) + if isinstance(step_val, dict): + stat = dict(step_val) + if isinstance(time_val, dict) and time_val.get("index") is not None: + stat["timestamp"] = time_val["index"] + entry[field] = stat + elif isinstance(time_val, dict): + entry[field] = dict(time_val) + merged.append(entry) + return merged + + +def extract_value_stats( + step_resp: ApiResponseType, + time_resp: ApiResponseType, + keys: List[str], +) -> List[Dict[str, Any]]: + """从并发的 step/time 响应中提取并合并 value stats。""" + step_list = step_resp.data if step_resp.ok and isinstance(step_resp.data, list) else [] + time_list = time_resp.data if time_resp.ok and isinstance(time_resp.data, list) else [] + return merge_value_stats(step_list, time_list, keys) + + +# --------------------------------------------------------------------------- +# 导出 CSV 流式解析 +# --------------------------------------------------------------------------- + + +@safe.decorator(message="Failed to download CSV") +def stream_export_csv( + client: "Client", + url: str, + keys: List[str], + rq: Optional[RangeQuery] = None, + timeout: int = 30, + x_key: str = "", +) -> Optional[Dict[str, List[Dict[str, Any]]]]: + """Stream-download wide-format export CSV and parse per-key rows. + + CSV layout (from ``POST /house/metrics/scalar/export``):: + + step, {exp}-{key1}_step, {exp}-{key1}_timestamp, + {exp}-{key2}_step, {exp}-{key2}_timestamp, … [, {exp}-{x}_step, {exp}-{x}_timestamp] + + One row per step; columns are interleaved as ``(value, timestamp)`` pairs. + The ``{exp}-{key}_step`` column actually holds the metric **value** despite + its name — the suffix is a House naming convention. + + When ``x_key`` is given (custom x axis), it is exported as one extra trailing + column. Because House pivots by ``GROUP BY step``, the x value sits on the + same row as every y value — alignment is applied in this single streaming + pass. + + **NaN 占位与契约说明**: + - 非自定义轴(纯 step 轴,``not x_key``):保持原有历史行为,单元格缺失或无法解析时跳过 + (不生成占位点,保证 ``head`` 计数为真实有效点数且不输出 NaN/null); + - 自定义 x 轴(``x_key`` 存在):无 timestamp 类过滤时,单元格缺失/NaN 以 ``float("nan")`` + 占位(自定义 x 缺失时 ``index`` 同样占位 NaN),保证多 key / x 轴之间按下标对齐;±Inf 照常保留; + - ``type="custom"`` 范围过滤:若指定了 ``start`` / ``end``,x 缺失为 NaN 的行不属于合法值域, + 予以过滤丢弃。 + """ + import csv + import time + from collections import deque + + resp = client._session.get(url, stream=True, timeout=timeout) + resp.raise_for_status() + resp.encoding = "utf-8" + lines = resp.iter_lines(decode_unicode=True) + next(lines, None) # skip header — column order is known from ``keys`` (+ optional trailing x) + + n_keys = len(keys) + # x 列定位:若 x_key 已在 keys 中,复用其对应列;若为额外追加列,占末尾 (value, timestamp) 槽位 + if not x_key: + x_value_col = None + elif x_key in keys: + x_value_col = 1 + keys.index(x_key) * 2 + else: + x_value_col = 1 + n_keys * 2 + tail_limit = rq.tail if rq is not None and rq.tail is not None else None + rows_per_key: List[Any] = [deque(maxlen=tail_limit) if tail_limit is not None else [] for _ in range(n_keys)] + + last_start_ts: Optional[int] = None + if rq is not None and rq.last is not None: + last_start_ts = int(time.time() * 1000) - rq.last + + # CSV is ``ORDER BY step`` — safe to break once step exceeds the range end. + # ``type="custom"`` filters on the (possibly non-monotonic) x value domain instead, + # so the step-order early break must be disabled there; head's count-based early + # stop below remains valid in every mode. + step_end_bound: Optional[int] = None + if rq is not None and rq.type not in ("timestamp", "custom") and last_start_ts is None and rq.end is not None: + # step/timestamp 语义下校验器已保证非负整值(float 存储,int() 截断安全) + step_end_bound = int(rq.end) + + head_limit = rq.head if rq is not None and rq.head is not None else None + _warned_missing_ts = False + + for row in csv.reader(lines): + if not row: + continue + try: + step = int(row[0]) + except (ValueError, IndexError): + continue + + if step_end_bound is not None and step > step_end_bound: + break + + # --- type="step":start 前的行直接跳过(行级提前过滤) --- + if rq is not None and rq.type == "step" and rq.start is not None and step < rq.start: + continue + + # --- 自定义 x 列:缺失/NaN → index 以 NaN 占位(不丢行,保持多 key 按下标对齐) --- + x_value: Optional[float] = None + if x_value_col is not None: + raw_x = row[x_value_col] if x_value_col < len(row) else "" + if raw_x: + try: + parsed_x = float(raw_x) + except ValueError: + parsed_x = None + if parsed_x is not None and not math.isnan(parsed_x): + x_value = parsed_x + + # --- type="custom":按 x 值域纯谓词过滤(不假设单调,无提前终止依据) --- + # 若指定了有界范围,缺失 x(x_value is None)的行不属于合法值域,予以过滤 + if rq is not None and rq.type == "custom": + if x_value is None: + if rq.start is not None or rq.end is not None: + continue + else: + if rq.start is not None and x_value < rq.start: + continue + if rq.end is not None and x_value > rq.end: + continue + + for i in range(n_keys): + vc = 1 + i * 2 # value column index + tc = 2 + i * 2 # timestamp column index + if vc >= len(row): + continue + raw_val = row[vc] + value: float + if not x_key: + # 纯 step 轴(非自定义):保持历史行为,缺失/无效时跳过当前 key 的该行 + if not raw_val: + continue + try: + value = float(raw_val) + except ValueError: + continue + else: + # 自定义 x 轴:缺失(空串)/ NaN / 无法解析 → float("nan") 占位;±Inf 照常保留 + try: + value = float(raw_val) if raw_val else float("nan") + except ValueError: + value = float("nan") + + ts: Optional[int] = None + if tc < len(row) and row[tc]: + try: + ts = int(row[tc]) + except ValueError: + pass + + if rq is not None: + # --- ``last`` mode: filter by timestamp >= (now - last) --- + if last_start_ts is not None: + if ts is None: + if not _warned_missing_ts: + console.warning("CSV row missing `timestamp` column.") + _warned_missing_ts = True + continue + if ts < last_start_ts: + continue + # --- timestamp range mode --- + elif rq.type == "timestamp": + if ts is None: + if not _warned_missing_ts: + console.warning("CSV row missing `timestamp` column.") + _warned_missing_ts = True + continue + if rq.start is not None and ts < rq.start: + continue + if rq.end is not None and ts > rq.end: + continue + + item: Dict[str, Any] = {"step": step, "value": value} + if x_value_col is not None: + # x 缺失/NaN 时以 NaN 占位,保持 index 槽位存在 + item["index"] = x_value if x_value is not None else float("nan") + if ts is not None: + item["timestamp"] = ts + rows_per_key[i].append(item) + + # head early-stop: all keys collected enough rows + if head_limit is not None and all(len(r) >= head_limit for r in rows_per_key): + break + + result: Dict[str, List[Dict[str, Any]]] = {} + for i, key in enumerate(keys): + rows = list(rows_per_key[i]) + if head_limit is not None: + rows = rows[:head_limit] + result[key] = rows + return result diff --git a/swanlab/api/helper/request.py b/swanlab/api/helper/request.py new file mode 100644 index 000000000..ae41457bc --- /dev/null +++ b/swanlab/api/helper/request.py @@ -0,0 +1,152 @@ +""" +@author: caddiesnew +@file: request.py +@time: 2026/9/3 +@description: House 查询请求辅助 — 请求体构建、预签名链接获取与媒体项构建 +""" + +from typing import Any, Dict, List, Optional + +from swanlab.api.base import BaseEntity +from swanlab.api.typings.metric import ApiMediaItemDataType + +# --------------------------------------------------------------------------- +# 请求体构建(纯函数) +# --------------------------------------------------------------------------- + + +def build_column_ref( + experiment_id: str, + created_at: int, + key: str, + root_pro_id: str = "", + root_exp_id: str = "", +) -> Dict[str, Any]: + # createdAt 为 House 查询的数据入库时间下界,必传 + ref: Dict[str, Any] = {"experimentId": experiment_id, "key": key, "createdAt": created_at} + if root_pro_id: + ref["rootProId"] = root_pro_id + if root_exp_id: + ref["rootExpId"] = root_exp_id + return ref + + +def build_scalar_payload( + project_id: str, + run_id: str, + created_at: int, + keys: List[str], + sample: int = 1500, + x_type: str = "step", + x_key: Optional[str] = None, + root_pro_id: str = "", + root_exp_id: str = "", +) -> Dict[str, Any]: + payload: Dict[str, Any] = { + "projectId": project_id, + "xType": x_type, + "range": [0, 0], + "columns": [build_column_ref(run_id, created_at, key, root_pro_id, root_exp_id) for key in keys], + "num": sample if sample <= 1500 else 1500, + } + # xKey != "" 时服务端自动把 range/join 语义切换为 x 值域;"range": [0, 0] 用不到 + if x_key: + payload["xKey"] = x_key + return payload + + +def build_media_payload( + project_id: str, + run_id: str, + created_at: int, + keys: List[str], + step: Optional[int] = None, + root_pro_id: str = "", + root_exp_id: str = "", +) -> Dict[str, Any]: + payload: Dict[str, Any] = { + "projectId": project_id, + "columns": [build_column_ref(run_id, created_at, key, root_pro_id, root_exp_id) for key in keys], + } + if step is not None: + payload["step"] = step + return payload + + +def build_export_payload( + project_id: str, + run_id: str, + created_at: int, + keys: List[str], + experiment_name: str = "", + root_pro_id: str = "", + root_exp_id: str = "", +) -> Dict[str, Any]: + """Build payload for ``POST /house/metrics/scalar/export``. + + ``experimentName`` is required by the API but only used for CSV column + headers — the actual query uses ``experimentId``. Falls back to ``run_id`` + when the real name is unavailable. + """ + exp_name = experiment_name or run_id + columns: List[Dict[str, Any]] = [] + for key in keys: + # createdAt 为 House 查询的数据入库时间下界,必传 + col: Dict[str, Any] = { + "experimentName": exp_name, + "experimentId": run_id, + "key": key, + "createdAt": created_at, + } + if root_pro_id: + col["rootProId"] = root_pro_id + if root_exp_id: + col["rootExpId"] = root_exp_id + columns.append(col) + return {"projectId": project_id, "columns": columns} + + +# --------------------------------------------------------------------------- +# 预签名链接获取与媒体项构建 +# --------------------------------------------------------------------------- + + +def fetch_presigned_urls(entity: BaseEntity, prefix: str, paths: List[str]) -> Dict[str, str]: + """批量获取预签名下载链接,返回 path → url 映射。""" + if not paths: + return {} + resp = entity._post("/resources/presigned/get", data={"prefix": prefix, "paths": paths}) + if not resp.ok or not isinstance(resp.data, dict): + return {} + urls = resp.data.get("urls") or [] + return dict(zip(paths, urls)) if urls else {} + + +def fetch_file_presigned_urls(entity: BaseEntity, paths: List[str]) -> Dict[str, str]: + """通过完整资源路径批量获取预签名下载链接,返回 path → url 映射。""" + if not paths: + return {} + resp = entity._post("/files/presigned/get", data={"paths": paths}) + if not resp.ok or not isinstance(resp.data, dict): + return {} + urls = resp.data.get("urls") or [] + return dict(zip(paths, urls)) if urls else {} + + +def build_media_items( + entry: Dict[str, Any], + url_map: Dict[str, str], +) -> List[ApiMediaItemDataType]: + """将单个 metric entry 的 data/more 合并为 items,注入预签名 url。""" + # 后端对"无数据"可能返回显式 null(dict.get 的默认值不生效),统一兜底为空列表 + paths = entry.get("data") or [] + mores = entry.get("more") or [] + items: List[ApiMediaItemDataType] = [] + for i, path in enumerate(paths): + item: ApiMediaItemDataType = {} + if path in url_map: + item["url"] = url_map[path] + if i < len(mores) and isinstance(mores[i], dict): + item.update(mores[i]) + items.append(item) + return items diff --git a/swanlab/api/utils.py b/swanlab/api/helper/utils.py similarity index 87% rename from swanlab/api/utils.py rename to swanlab/api/helper/utils.py index d35c91aca..925bc346c 100644 --- a/swanlab/api/utils.py +++ b/swanlab/api/helper/utils.py @@ -16,6 +16,7 @@ ApiFilterStableKeyLiteral, ApiMetricAllTypeLiteral, ApiMetricLogLevelLiteral, + ApiMetricXAxisLiteral, ApiSidebarLiteral, ApiSortOrderLiteral, ApiVisibilityLiteral, @@ -99,6 +100,7 @@ def validate_non_empty_string(value: str, *, label: str) -> None: # 指标相关校验常量 _VALID_METRIC_ALL_TYPES = frozenset(get_args(ApiMetricAllTypeLiteral)) _VALID_METRIC_LOG_LEVELS = frozenset(get_args(ApiMetricLogLevelLiteral)) +_VALID_X_AXIS_LITERALS = frozenset(get_args(ApiMetricXAxisLiteral)) def _check_required(item: Dict[str, Any], keys: Set[str]) -> None: @@ -256,3 +258,30 @@ def validate_metric_keys(keys: List[str]) -> None: """校验 metric keys 列表的合法性。""" if not isinstance(keys, list) or not keys or any(not isinstance(key, str) or not key.strip() for key in keys): raise ValueError("keys must be a non-empty list of non-empty strings") + + +def validate_x_axis(x_axis: str, *, metric_type: str = "SCALAR") -> str: + """校验 x_axis 参数并原样返回。 + + 内置轴字面量(step/time/relative_time)按已知集合校验;其余非空字符串放行为自定义 + x 列 key。非 SCALAR 查询(MEDIA/LOG)没有 x 轴概念,只允许 step,其余取值尽早报错。 + """ + if not isinstance(x_axis, str) or not x_axis.strip(): + raise ValueError("x_axis must be a non-empty string") + if x_axis == "auto": + raise ValueError( + "x_axis='auto' is not supported; pass 'step', a built-in axis ('time', 'relative_time'), " + "or an explicit custom column key." + ) + if x_axis == "timestamp": + raise ValueError( + "Invalid x_axis 'timestamp'. Built-in time axis is 'time' (or 'relative_time'); " + "'timestamp' is only supported as a range_query type." + ) + if x_axis in _VALID_X_AXIS_LITERALS: + if metric_type != "SCALAR" and x_axis != "step": + raise ValueError(f"x_axis={x_axis!r} is only supported for SCALAR metrics") + return x_axis + if metric_type != "SCALAR": + raise ValueError(f"custom x_axis key {x_axis!r} is only supported for SCALAR metrics") + return x_axis diff --git a/swanlab/api/metric.py b/swanlab/api/metric.py index 8596e9408..93a309f12 100644 --- a/swanlab/api/metric.py +++ b/swanlab/api/metric.py @@ -7,9 +7,29 @@ from __future__ import annotations -from typing import TYPE_CHECKING, Any, Callable, Dict, Iterator, List, Optional +from typing import Any, Dict, Iterator, List, Optional, cast from swanlab.api.base import ApiClientContext, BaseEntity +from swanlab.api.helper import ( + SCALAR_STATISTIC_FIELDS, + align_entries_by_key, + axis_request_params, + build_export_payload, + build_media_items, + build_media_payload, + build_scalar_payload, + builtin_x_axis, + extract_first, + extract_value_stats, + fetch_file_presigned_urls, + fetch_presigned_urls, + get_properties, + stream_export_csv, + validate_metric_keys, + validate_metric_log_level, + validate_metric_type, + validate_x_axis, +) from swanlab.api.typings import ApiColumnCsvExportType, ApiResponseType from swanlab.api.typings.common import ( MAX_CONCURRENT_COUNT, @@ -20,180 +40,16 @@ ) from swanlab.api.typings.metric import ( ApiLogSeriesType, - ApiMediaItemDataType, ApiMediaSeriesType, + ApiMetricXAxisType, ApiScalarSeriesType, ) -from swanlab.api.utils import get_properties, validate_metric_keys, validate_metric_log_level, validate_metric_type -from swanlab.sdk.internal.pkg import console, safe +from swanlab.sdk.internal.pkg import console from swanlab.sdk.internal.pkg.executor import SafeThreadPoolExecutor -if TYPE_CHECKING: - from swanlab.sdk.internal.pkg.client import Client - -_SCALAR_STATISTIC_FIELDS = ("min", "max", "avg", "median", "latest") _METRIC_SHARED_KEYS = frozenset({"project_id", "run_id", "metric_type"}) -def _align_entries_by_key(entries: List[Dict[str, Any]]) -> Dict[str, Dict[str, Any]]: - """将后端返回的列表按 ``key`` 字段映射为 dict,应对后端可能省略列或乱序返回。""" - indexed: Dict[str, Dict[str, Any]] = {} - for entry in entries: - if not isinstance(entry, dict): - continue - key = entry.get("key") - if isinstance(key, str) and key: - indexed[key] = entry - return indexed - - -def _merge_value_stats( - step_list: List[Dict[str, Any]], - time_list: List[Dict[str, Any]], - keys: List[str], -) -> List[Dict[str, Any]]: - """合并 step/time 两种 x_type 的 value stat 响应为 per-key 统计字典。""" - step_by_key = _align_entries_by_key(step_list) - time_by_key = _align_entries_by_key(time_list) - merged: List[Dict[str, Any]] = [] - for key in keys: - step_entry = step_by_key.get(key, {}) - time_entry = time_by_key.get(key, {}) - entry: Dict[str, Any] = {} - for field in _SCALAR_STATISTIC_FIELDS: - step_val = step_entry.get(field) - time_val = time_entry.get(field) - if isinstance(step_val, dict): - stat = dict(step_val) - if isinstance(time_val, dict) and time_val.get("index") is not None: - stat["timestamp"] = time_val["index"] - entry[field] = stat - elif isinstance(time_val, dict): - entry[field] = dict(time_val) - merged.append(entry) - return merged - - -@safe.decorator(message="Failed to download CSV") -def _stream_export_csv( - client: "Client", - url: str, - keys: List[str], - rq: Optional[RangeQuery] = None, - timeout: int = 30, -) -> Optional[Dict[str, List[Dict[str, Any]]]]: - """Stream-download wide-format export CSV and parse per-key rows. - - CSV layout (from ``POST /house/metrics/scalar/export``):: - - step, {exp}-{key1}_step, {exp}-{key1}_timestamp, - {exp}-{key2}_step, {exp}-{key2}_timestamp, … - - One row per step; columns are interleaved as ``(value, timestamp)`` pairs. - The ``{exp}-{key}_step`` column actually holds the metric **value** despite - its name — the suffix is a House naming convention. - """ - import csv - import time - from collections import deque - - resp = client._session.get(url, stream=True, timeout=timeout) - resp.raise_for_status() - resp.encoding = "utf-8" - lines = resp.iter_lines(decode_unicode=True) - next(lines, None) # skip header — column order is known from ``keys`` - - n_keys = len(keys) - tail_limit = rq.tail if rq is not None and rq.tail is not None else None - rows_per_key: List[Any] = [deque(maxlen=tail_limit) if tail_limit is not None else [] for _ in range(n_keys)] - - last_start_ts: Optional[int] = None - if rq is not None and rq.last is not None: - last_start_ts = int(time.time() * 1000) - rq.last - - # CSV is ``ORDER BY step`` — safe to break once step exceeds the range end. - step_end_bound: Optional[int] = None - if rq is not None and rq.type != "timestamp" and last_start_ts is None and rq.end is not None: - step_end_bound = rq.end - - head_limit = rq.head if rq is not None and rq.head is not None else None - _warned_missing_ts = False - - for row in csv.reader(lines): - if not row: - continue - try: - step = int(row[0]) - except (ValueError, IndexError): - continue - - if step_end_bound is not None and step > step_end_bound: - break - - for i in range(n_keys): - vc = 1 + i * 2 # value column index - tc = 2 + i * 2 # timestamp column index - if vc >= len(row): - continue - raw_val = row[vc] - if not raw_val: - continue - try: - value = float(raw_val) - except ValueError: - continue - - ts: Optional[int] = None - if tc < len(row) and row[tc]: - try: - ts = int(row[tc]) - except ValueError: - pass - - if rq is not None: - # --- ``last`` mode: filter by timestamp >= (now - last) --- - if last_start_ts is not None: - if ts is None: - if not _warned_missing_ts: - console.warning("CSV row missing `timestamp` column.") - _warned_missing_ts = True - continue - if ts < last_start_ts: - continue - # --- timestamp range mode --- - elif rq.type == "timestamp": - if ts is None: - if not _warned_missing_ts: - console.warning("CSV row missing `timestamp` column.") - _warned_missing_ts = True - continue - if rq.start is not None and ts < rq.start: - continue - if rq.end is not None and ts > rq.end: - continue - # --- step range mode (end handled by step_end_bound break) --- - else: - if rq.start is not None and step < rq.start: - continue - - item: Dict[str, Any] = {"step": step, "value": value} - if ts is not None: - item["timestamp"] = ts - rows_per_key[i].append(item) - - # head early-stop: all keys collected enough rows - if head_limit is not None and all(len(r) >= head_limit for r in rows_per_key): - break - - result: Dict[str, List[Dict[str, Any]]] = {} - for i, key in enumerate(keys): - rows = list(rows_per_key[i]) - if head_limit is not None: - rows = rows[:head_limit] - result[key] = rows - return result - - class Metric(BaseEntity): """ 表示一个 SwanLab 指标列(非单个数值,而是一组序列)。 @@ -220,11 +76,13 @@ def __init__( root_pro_id: str = "", root_exp_id: str = "", experiment_name: str = "", + x_axis: str = "step", ) -> None: super().__init__(ctx) validate_metric_type(metric_type, key) if metric_type == "LOG": validate_metric_log_level(log_level) + validate_x_axis(x_axis, metric_type=metric_type) self._project_id = project_id self._run_id = run_id self._key = key @@ -242,6 +100,7 @@ def __init__( self._root_exp_id = root_exp_id self._created_at = created_at self._experiment_name = experiment_name + self._x_axis = x_axis # 类型 → 加载方法 的分发表,新增类型只需在此注册 _FETCH_DISPATCH = { @@ -308,99 +167,6 @@ def step(self) -> Optional[int]: # 请求辅助函数 # ------------------------------------------------------------------ - @staticmethod - def _extract_first(resp: ApiResponseType) -> Optional[Dict[str, Any]]: - """从列表型 API 响应中提取第一个元素,失败返回 None。""" - if resp.ok and isinstance(resp.data, list) and resp.data: - return resp.data[0] - return None - - @staticmethod - def _build_column_ref( - experiment_id: str, - created_at: int, - key: str, - root_pro_id: str = "", - root_exp_id: str = "", - ) -> Dict[str, Any]: - # createdAt 为 House 查询的数据入库时间下界,必传 - ref: Dict[str, Any] = {"experimentId": experiment_id, "key": key, "createdAt": created_at} - if root_pro_id: - ref["rootProId"] = root_pro_id - if root_exp_id: - ref["rootExpId"] = root_exp_id - return ref - - @staticmethod - def _build_scalar_payload( - project_id: str, - run_id: str, - created_at: int, - keys: List[str], - sample: int = 1500, - x_type: str = "step", - root_pro_id: str = "", - root_exp_id: str = "", - ) -> Dict[str, Any]: - return { - "projectId": project_id, - "xType": x_type, - "range": [0, 0], - "columns": [Metric._build_column_ref(run_id, created_at, key, root_pro_id, root_exp_id) for key in keys], - "num": sample if sample <= 1500 else 1500, - } - - @staticmethod - def _build_media_payload( - project_id: str, - run_id: str, - created_at: int, - keys: List[str], - step: Optional[int] = None, - root_pro_id: str = "", - root_exp_id: str = "", - ) -> Dict[str, Any]: - payload: Dict[str, Any] = { - "projectId": project_id, - "columns": [Metric._build_column_ref(run_id, created_at, key, root_pro_id, root_exp_id) for key in keys], - } - if step is not None: - payload["step"] = step - return payload - - @staticmethod - def _build_export_payload( - project_id: str, - run_id: str, - created_at: int, - keys: List[str], - experiment_name: str = "", - root_pro_id: str = "", - root_exp_id: str = "", - ) -> Dict[str, Any]: - """Build payload for ``POST /house/metrics/scalar/export``. - - ``experimentName`` is required by the API but only used for CSV column - headers — the actual query uses ``experimentId``. Falls back to ``run_id`` - when the real name is unavailable. - """ - exp_name = experiment_name or run_id - columns: List[Dict[str, Any]] = [] - for key in keys: - # createdAt 为 House 查询的数据入库时间下界,必传 - col: Dict[str, Any] = { - "experimentName": exp_name, - "experimentId": run_id, - "key": key, - "createdAt": created_at, - } - if root_pro_id: - col["rootProId"] = root_pro_id - if root_exp_id: - col["rootExpId"] = root_exp_id - columns.append(col) - return {"projectId": project_id, "columns": columns} - def _build_log_params(self) -> Dict[str, Any]: params: Dict[str, Any] = { "projectId": self.project_id, @@ -422,92 +188,39 @@ def _build_log_params(self) -> Dict[str, Any]: # ------------------------------------------------------------------ def _fetch_scalar(self) -> ApiScalarSeriesType: - res = ApiScalarSeriesType(projectId=self.project_id, experimentId=self.run_id, key=self.key) + """委托给单 key Metrics 获取标量数据(统一走轴发现、采样/全量 CSV 及回显)。""" + metrics_obj = Metrics( + self._ctx, + project_id=self._project_id, + run_id=self._run_id, + keys=[self.key], + metric_type="SCALAR", + sample=self._sample, + ignore_timestamp=self._ignore_timestamp, + all=self._all, + x_axis=self._x_axis, + root_pro_id=self._root_pro_id, + root_exp_id=self._root_exp_id, + created_at=self._created_at, + experiment_name=self._experiment_name, + ) + batch = metrics_obj._ensure_batch() + if batch and batch[0]._data: + return cast(ApiScalarSeriesType, batch[0]._data) + res = ApiScalarSeriesType(projectId=self.project_id, experimentId=self.run_id, key=self.key, metrics=[]) if self._root_pro_id: res["rootProId"] = self._root_pro_id if self._root_exp_id: res["rootExpId"] = self._root_exp_id - payload = self._build_scalar_payload( - self.project_id, - self.run_id, - self._created_at, - [self.key], - self._sample, - root_pro_id=self._root_pro_id, - root_exp_id=self._root_exp_id, - ) - - # 1. 获取折线数据 — 使用 key-indexed lookup 保证对齐 - scalar_resp = self._post("/house/metrics/scalar", data=payload) - if scalar_resp.ok and isinstance(scalar_resp.data, list): - scalar_by_key = _align_entries_by_key(scalar_resp.data) - res["metrics"] = scalar_by_key.get(self.key, {}).get("metrics", []) - if not res.get("metrics"): - return res - - # 2. 获取统计值 — step/time 并发,key-indexed 合并 - step_payload = {**payload, "xType": "step"} - time_payload = {**payload, "xType": "timestamp"} - step_resp, time_resp = self._concurrent_request( - [ - (self._post, "/house/metrics/scalar/value", {"data": step_payload}), - (self._post, "/house/metrics/scalar/value", {"data": time_payload}), - ] - ) - step_list = step_resp.data if step_resp.ok and isinstance(step_resp.data, list) else [] - time_list = time_resp.data if time_resp.ok and isinstance(time_resp.data, list) else [] - value_list = _merge_value_stats(step_list, time_list, [self.key]) - if value_list: - for field in _SCALAR_STATISTIC_FIELDS: - val = value_list[0].get(field) - if val: - res[field] = val return res - @staticmethod - def _fetch_presigned_urls(entity: BaseEntity, prefix: str, paths: List[str]) -> Dict[str, str]: - """批量获取预签名下载链接,返回 path → url 映射。""" - if not paths: - return {} - resp = entity._post("/resources/presigned/get", data={"prefix": prefix, "paths": paths}) - if not resp.ok or not isinstance(resp.data, dict): - return {} - urls = resp.data.get("urls") or [] - return dict(zip(paths, urls)) if urls else {} - - @staticmethod - def _fetch_file_presigned_urls(entity: BaseEntity, paths: List[str]) -> Dict[str, str]: - """通过完整资源路径批量获取预签名下载链接,返回 path → url 映射。""" - if not paths: - return {} - resp = entity._post("/files/presigned/get", data={"paths": paths}) - if not resp.ok or not isinstance(resp.data, dict): - return {} - urls = resp.data.get("urls") or [] - return dict(zip(paths, urls)) if urls else {} - - @staticmethod - def _build_media_items( - entry: Dict[str, Any], - url_map: Dict[str, str], - ) -> List[ApiMediaItemDataType]: - """将单个 metric entry 的 data/more 合并为 items,注入预签名 url。""" - # 后端对"无数据"可能返回显式 null(dict.get 的默认值不生效),统一兜底为空列表 - paths = entry.get("data") or [] - mores = entry.get("more") or [] - items: List[ApiMediaItemDataType] = [] - for i, path in enumerate(paths): - item: ApiMediaItemDataType = {} - if path in url_map: - item["url"] = url_map[path] - if i < len(mores) and isinstance(mores[i], dict): - item.update(mores[i]) - items.append(item) - return items - def _fetch_media(self) -> ApiMediaSeriesType: res = ApiMediaSeriesType(projectId=self.project_id, experimentId=self.run_id, key=self.key) - payload = self._build_media_payload( + if self._root_pro_id: + res["rootProId"] = self._root_pro_id + if self._root_exp_id: + res["rootExpId"] = self._root_exp_id + payload = build_media_payload( self.project_id, self.run_id, self._created_at, @@ -536,18 +249,22 @@ def _fetch_media(self) -> ApiMediaSeriesType: prefix = f"{self.project_id}/{self.run_id}" all_paths = metric_entry.get("data") or [] - url_map = self._fetch_presigned_urls(self, prefix, all_paths) if all_paths else {} + url_map = fetch_presigned_urls(self, prefix, all_paths) if all_paths else {} if all_paths: console.debug( f"Media fetched: run_id[{self.run_id}], key[{self.key}] - {len(all_paths)} items, requesting presigned urls..." ) - items = self._build_media_items(metric_entry, url_map) + items = build_media_items(metric_entry, url_map) res["metrics"] = [{"index": data.get("step") or 0, "items": items}] return res def _fetch_media_all(self) -> ApiMediaSeriesType: res = ApiMediaSeriesType(projectId=self.project_id, experimentId=self.run_id, key=self.key) - payload = self._build_media_payload( + if self._root_pro_id: + res["rootProId"] = self._root_pro_id + if self._root_exp_id: + res["rootExpId"] = self._root_exp_id + payload = build_media_payload( self.project_id, self.run_id, self._created_at, @@ -556,7 +273,7 @@ def _fetch_media_all(self) -> ApiMediaSeriesType: root_exp_id=self._root_exp_id, ) raw_resp = self._post("/house/metrics/f_media", data=payload) - raw_data = self._extract_first(raw_resp) + raw_data = extract_first(raw_resp) if raw_data is None: return res @@ -564,13 +281,13 @@ def _fetch_media_all(self) -> ApiMediaSeriesType: # metrics 可能为 null 或含 None 条目,过滤后再展开 entries = [e for e in (raw_data.get("metrics") or []) if isinstance(e, dict)] all_paths = [p for entry in entries for p in (entry.get("data") or [])] - url_map = self._fetch_presigned_urls(self, prefix, all_paths) if all_paths else {} + url_map = fetch_presigned_urls(self, prefix, all_paths) if all_paths else {} if all_paths: console.debug( f"Media fetched (all): run_id[{self.run_id}], key[{self.key}] - {len(all_paths)} items, requesting presigned urls..." ) res["metrics"] = [ - {"index": entry.get("index", 0), "items": self._build_media_items(entry, url_map)} for entry in entries + {"index": entry.get("index", 0), "items": build_media_items(entry, url_map)} for entry in entries ] return res @@ -598,7 +315,7 @@ def export_csv(self) -> ApiResponseType: """ if self.metric_type != "SCALAR": return ApiResponseType(ok=False, errmsg="export_csv() only support SCALAR metric_type", data=None) - payload = Metric._build_export_payload( + payload = build_export_payload( self._project_id, self._run_id, self._created_at, @@ -613,7 +330,7 @@ def export_csv(self) -> ApiResponseType: cos_key = resp.data.get("cosKey", "") if isinstance(resp.data, dict) else "" if not cos_key: return ApiResponseType(ok=False, errmsg="Invalid response format: missing cosKey", data=None) - url_map = Metric._fetch_file_presigned_urls(self, [cos_key]) + url_map = fetch_file_presigned_urls(self, [cos_key]) url = url_map.get(cos_key, "") if not url: return ApiResponseType(ok=False, errmsg="Failed to get presigned download URL", data=None) @@ -641,7 +358,7 @@ def export_logs(self, start: int = 0, rows: int = 500_000) -> ApiResponseType: cos_key = resp.data.get("cosKey", "") if not cos_key: return ApiResponseType(ok=False, errmsg="Invalid response format: missing cosKey", data=None) - url_map = Metric._fetch_file_presigned_urls(self, [cos_key]) + url_map = fetch_file_presigned_urls(self, [cos_key]) url = url_map.get(cos_key, "") if not url: return ApiResponseType(ok=False, errmsg="Failed to get presigned download URL", data=None) @@ -650,12 +367,19 @@ def export_logs(self, start: int = 0, rows: int = 500_000) -> ApiResponseType: def json(self) -> Dict[str, Any]: result = get_properties(self) data = self._ensure_data() + if "rootProId" in data: + result["rootProId"] = data["rootProId"] + if "rootExpId" in data: + result["rootExpId"] = data["rootExpId"] if self._metric_type == "SCALAR": if "url" in data: result.pop("metrics", None) result["url"] = data["url"] - for field in _SCALAR_STATISTIC_FIELDS: + # 回显 metrics[].index 的实际轴语义(自定义 x / 内置轴 / 回落 step) + if "xAxis" in data: + result["xAxis"] = data["xAxis"] + for field in SCALAR_STATISTIC_FIELDS: val = data.get(field) if val: result[field] = val @@ -722,6 +446,7 @@ def __init__( root_exp_id: str = "", created_at: int, experiment_name: str = "", + x_axis: str = "step", ) -> None: super().__init__(ctx) validate_metric_keys(keys) @@ -730,6 +455,19 @@ def __init__( raise ValueError("Metrics does not support LOG metric_type, use Experiment.logs() instead") if range_query is not None and metric_type != "SCALAR": raise ValueError("range_query is only supported for SCALAR metric_type") + validate_x_axis(x_axis, metric_type=metric_type) + if (range_query is not None or all) and x_axis in ("time", "relative_time"): + raise ValueError( + f"x_axis={x_axis!r} is not supported in CSV mode (all=True or range_query); " + "use the sampled mode (default) or a custom x column key instead" + ) + if range_query is not None and range_query.type == "custom": + builtin = builtin_x_axis(x_axis) + if builtin is not None and builtin.get("type") != "CUSTOM": + raise ValueError( + "range_query type='custom' filters on custom x-axis values and requires a custom x axis; " + f"got x_axis={x_axis!r}. Pass x_axis= instead" + ) self._project_id = project_id self._run_id = run_id # 去重,保持插入顺序 @@ -743,6 +481,7 @@ def __init__( self._root_exp_id = root_exp_id self._created_at = created_at self._experiment_name = experiment_name + self._x_axis = x_axis self._page_info: Dict[str, Any] = { "keys": keys, "metricType": metric_type, @@ -768,6 +507,9 @@ def _ensure_batch(self) -> List[Metric]: self._cached_list = list(self._fetch_batch()) return self._cached_list + def _ensure_data(self) -> Any: + return self._ensure_batch() + def __iter__(self) -> Iterator[Metric]: yield from self._ensure_batch() @@ -785,9 +527,9 @@ def _fetch_batch(self) -> Iterator[Metric]: """根据 metric_type 和模式分发到具体的获取方法。""" if self._metric_type == "SCALAR": if self._range_query is not None or self._all: - data_list = self._fetch_scalar_csv(self._keys) + data_list = self._fetch_scalar_csv_data() else: - data_list = self._batch_keys(self._fetch_scalar_lines) + data_list = self._fetch_scalar_sampled_data() else: # media 后端已支持 columns 批量,无需分批 if self._all: @@ -798,15 +540,48 @@ def _fetch_batch(self) -> Iterator[Metric]: for data in data_list: yield self._build_metric(data.get("key", ""), data) - def _batch_keys(self, fetch_fn: Callable[[List[str]], List[Dict[str, Any]]]) -> List[Dict[str, Any]]: + # ------------------------------------------------------------------ + # X 轴解析与数据获取 + # ------------------------------------------------------------------ + + def _resolve_axis(self) -> ApiMetricXAxisType: + """解析当前查询的 X 轴(显式内置轴或自定义列 key)。""" + builtin = builtin_x_axis(self._x_axis) + if builtin is not None: + return builtin + return {"type": "CUSTOM", "key": self._x_axis} + + def _fetch_scalar_sampled_data(self) -> List[Dict[str, Any]]: + """采样路径:单次请求统一按 self._x_axis 请求 House,按 keys 顺序回显。""" + axis = self._resolve_axis() + x_type, x_key = axis_request_params(axis) + results = self._batch_keys(self._keys, x_type=x_type, x_key=x_key) + for r in results: + r["xAxis"] = axis + return results + + def _fetch_scalar_csv_data(self) -> List[Dict[str, Any]]: + """CSV 全量路径:导出并提取单轴数据。""" + axis = self._resolve_axis() + x_key = axis.get("key", "") if axis.get("type") == "CUSTOM" else "" + results = self._fetch_scalar_csv(self._keys, x_key=x_key) + for r in results: + r["xAxis"] = axis + return results + + def _batch_keys( + self, + keys: List[str], + x_type: str = "step", + x_key: str = "", + ) -> List[Dict[str, Any]]: """将 keys 按 ``_BATCH_SIZE`` 分批;单批直接执行,多批并发。""" - keys = self._keys if len(keys) <= self._BATCH_SIZE: - return fetch_fn(keys) + return self._fetch_scalar_lines(keys, x_type=x_type, x_key=x_key) chunks = [keys[i : i + self._BATCH_SIZE] for i in range(0, len(keys), self._BATCH_SIZE)] with SafeThreadPoolExecutor(max_workers=min(len(chunks), self._BATCH_SIZE)) as pool: - futures = [pool.submit(fetch_fn, chunk) for chunk in chunks] + futures = [pool.submit(self._fetch_scalar_lines, chunk, x_type, x_key) for chunk in chunks] results: List[Dict[str, Any]] = [] for f in futures: results.extend(f.result()) @@ -840,7 +615,20 @@ def _build_metric(self, key: str, data: Dict[str, Any]) -> Metric: def _empty_scalar_results(self, keys: List[str]) -> List[Dict[str, Any]]: """返回 per-key 空结果列表(用于 early return)。""" - return [{"projectId": self._project_id, "experimentId": self._run_id, "key": k, "metrics": []} for k in keys] + results: List[Dict[str, Any]] = [] + for k in keys: + data: Dict[str, Any] = { + "projectId": self._project_id, + "experimentId": self._run_id, + "key": k, + "metrics": [], + } + if self._root_pro_id: + data["rootProId"] = self._root_pro_id + if self._root_exp_id: + data["rootExpId"] = self._root_exp_id + results.append(data) + return results def _build_value_stats_requests(self, keys: List[str]) -> List[tuple]: """构建 step/time value stats 的并发请求列表(2 路并发)。""" @@ -850,7 +638,7 @@ def _build_value_stats_requests(self, keys: List[str]) -> List[tuple]: self._post, value_path, { - "data": Metric._build_scalar_payload( + "data": build_scalar_payload( self._project_id, self._run_id, self._created_at, @@ -865,26 +653,16 @@ def _build_value_stats_requests(self, keys: List[str]) -> List[tuple]: for x_type in ("step", "timestamp") ] - @staticmethod - def _extract_value_stats( - step_resp: ApiResponseType, - time_resp: ApiResponseType, - keys: List[str], - ) -> List[Dict[str, Any]]: - """从并发的 step/time 响应中提取并合并 value stats。""" - step_list = step_resp.data if step_resp.ok and isinstance(step_resp.data, list) else [] - time_list = time_resp.data if time_resp.ok and isinstance(time_resp.data, list) else [] - return _merge_value_stats(step_list, time_list, keys) - # ------------------------------------------------------------------ # Scalar: 折线数据 + 统计值 (后端 columns 批量) # ------------------------------------------------------------------ - def _fetch_scalar_lines(self, keys: List[str]) -> List[Dict[str, Any]]: + def _fetch_scalar_lines(self, keys: List[str], x_type: str = "step", x_key: str = "") -> List[Dict[str, Any]]: """获取标量折线数据 + step/time 统计值,3 路并发。 后端 ``POST /house/metrics/scalar`` 和 ``/scalar/value`` 的 ``columns`` 数组天然支持多 key,此处将 keys 打包为一个批量请求。 + 同一分组内的 keys 共享一个轴(xKey / xType);stats 请求不支持 xKey,不带。 """ # 3 路并发:折线数据 + step 统计 + time 统计 requests: List[tuple] = [ @@ -892,12 +670,14 @@ def _fetch_scalar_lines(self, keys: List[str]) -> List[Dict[str, Any]]: self._post, "/house/metrics/scalar", { - "data": Metric._build_scalar_payload( + "data": build_scalar_payload( self._project_id, self._run_id, self._created_at, keys, self._sample, + x_type=x_type, + x_key=x_key or None, root_pro_id=self._root_pro_id, root_exp_id=self._root_exp_id, ) @@ -909,9 +689,9 @@ def _fetch_scalar_lines(self, keys: List[str]) -> List[Dict[str, Any]]: scalar_resp, step_resp, time_resp = self._concurrent_request(requests) scalar_list = scalar_resp.data if scalar_resp.ok and isinstance(scalar_resp.data, list) else [] - scalar_by_key = _align_entries_by_key(scalar_list) + scalar_by_key = align_entries_by_key(scalar_list) metrics_by_key: Dict[str, Any] = {key: scalar_by_key.get(key, {}).get("metrics", []) for key in keys} - value_list = self._extract_value_stats(step_resp, time_resp, keys) + value_list = extract_value_stats(step_resp, time_resp, keys) value_by_key: Dict[str, Dict[str, Any]] = {keys[i]: v for i, v in enumerate(value_list)} return self._build_scalar_results(keys, metrics_by_key, value_by_key) @@ -920,39 +700,45 @@ def _fetch_scalar_lines(self, keys: List[str]) -> List[Dict[str, Any]]: # Scalar: CSV 全量下载 + 统计值 (range_query 或 all 模式) # ------------------------------------------------------------------ - def _fetch_scalar_csv(self, keys: List[str]) -> List[Dict[str, Any]]: + def _fetch_scalar_csv(self, keys: List[str], x_key: str = "") -> List[Dict[str, Any]]: """CSV 全量下载 + value stats,使用 House 批量导出接口。 - 将 keys 按 ``_CSV_KEY_BATCH_SIZE``(16)个一组分批,4 线程并发获取。 + 将 keys 按 ``MAX_METRIC_KEY_BATCH_SIZE``(16)个一组分批,4 线程并发获取。 + 自定义 x 轴下 ``x_key`` 作为额外的普通导出列追加,占用一个分批槽位 + (每批 y keys 上限降为 15)。 通过 ``POST /house/metrics/scalar/export`` 一次性导出每批 key 到 CSV 文件, 再通过 ``/files/presigned/get`` 获取预签名下载链接。 value stats 批量获取,与导出请求并发执行。 """ - if len(keys) <= MAX_METRIC_KEY_BATCH_SIZE: - return self._fetch_scalar_csv_batch(keys) + chunk_size = MAX_METRIC_KEY_BATCH_SIZE - 1 if x_key else MAX_METRIC_KEY_BATCH_SIZE + if len(keys) <= chunk_size: + return self._fetch_scalar_csv_batch(keys, x_key=x_key) - chunks = [keys[i : i + MAX_METRIC_KEY_BATCH_SIZE] for i in range(0, len(keys), MAX_METRIC_KEY_BATCH_SIZE)] + chunks = [keys[i : i + chunk_size] for i in range(0, len(keys), chunk_size)] with SafeThreadPoolExecutor(max_workers=MAX_CONCURRENT_COUNT) as pool: - futures = [pool.submit(self._fetch_scalar_csv_batch, chunk) for chunk in chunks] + futures = [pool.submit(self._fetch_scalar_csv_batch, chunk, x_key) for chunk in chunks] results: List[Dict[str, Any]] = [] for f in futures: results.extend(f.result()) return results - def _fetch_scalar_csv_batch(self, keys: List[str]) -> List[Dict[str, Any]]: - """单批 CSV 导出 + value stats(≤16 keys)。 + def _fetch_scalar_csv_batch(self, keys: List[str], x_key: str = "") -> List[Dict[str, Any]]: + """单批 CSV 导出 + value stats(≤16 列;自定义 x 时 ≤15 y keys + 1 x 列)。 通过 ``POST /house/metrics/scalar/export`` 一次性导出所有 key 到一个 CSV 文件, 再通过 ``/files/presigned/get`` 获取预签名下载链接。 - value stats 批量获取,与导出请求并发执行。 + value stats 批量获取,与导出请求并发执行;stats 请求只含 y keys,不含 x 列。 """ # 并发:step 统计 + time 统计 + 批量 CSV 导出 requests: List[tuple] = self._build_value_stats_requests(keys) - export_payload = Metric._build_export_payload( + export_columns = list(keys) + if x_key and x_key not in keys: + export_columns.append(x_key) + export_payload = build_export_payload( self._project_id, self._run_id, self._created_at, - keys, + export_columns, experiment_name=self._experiment_name, root_pro_id=self._root_pro_id, root_exp_id=self._root_exp_id, @@ -963,7 +749,7 @@ def _fetch_scalar_csv_batch(self, keys: List[str]) -> List[Dict[str, Any]]: step_resp, time_resp = all_resps[0], all_resps[1] export_resp = all_resps[2] - value_list = self._extract_value_stats(step_resp, time_resp, keys) + value_list = extract_value_stats(step_resp, time_resp, keys) value_by_key: Dict[str, Dict[str, Any]] = {keys[i]: v for i, v in enumerate(value_list)} # cosKey → presigned URL → download → parse per-key rows @@ -972,10 +758,10 @@ def _fetch_scalar_csv_batch(self, keys: List[str]) -> List[Dict[str, Any]]: if export_resp.ok and isinstance(export_resp.data, dict): cos_key = export_resp.data.get("cosKey", "") if cos_key: - url_map = Metric._fetch_file_presigned_urls(self, [cos_key]) + url_map = fetch_file_presigned_urls(self, [cos_key]) url = url_map.get(cos_key, "") if url: - parsed = _stream_export_csv(self._ctx.client, url, keys, self._range_query) + parsed = stream_export_csv(self._ctx.client, url, keys, self._range_query, x_key=x_key) if parsed: metrics_by_key = parsed @@ -1003,6 +789,10 @@ def _build_scalar_results( "key": key, "metrics": metrics_by_key.get(key, []), } + if self._root_pro_id: + data["rootProId"] = self._root_pro_id + if self._root_exp_id: + data["rootExpId"] = self._root_exp_id stats = value_by_key.get(key, {}) if stats: data.update(stats) @@ -1015,7 +805,7 @@ def _build_scalar_results( def _fetch_media_data(self) -> List[Dict[str, Any]]: """获取媒体数据(单步),后端 columns 批量一次返回。""" - payload = Metric._build_media_payload( + payload = build_media_payload( self._project_id, self._run_id, self._created_at, @@ -1038,7 +828,7 @@ def _fetch_media_data(self) -> List[Dict[str, Any]]: prefix = f"{self._project_id}/{self._run_id}" all_paths = [p for entry in metrics_raw for p in (entry.get("data") or [])] - url_map = Metric._fetch_presigned_urls(self, prefix, all_paths) if all_paths else {} + url_map = fetch_presigned_urls(self, prefix, all_paths) if all_paths else {} if all_paths: console.debug( f"Media fetched: run_id[{self._run_id}] - {len(all_paths)} items across {len(self._keys)} keys, requesting presigned urls..." @@ -1055,16 +845,20 @@ def _fetch_media_data(self) -> List[Dict[str, Any]]: "step": current_step, "metrics": [], } + if self._root_pro_id: + data["rootProId"] = self._root_pro_id + if self._root_exp_id: + data["rootExpId"] = self._root_exp_id entry = key_to_entry.get(key) if entry: - items = Metric._build_media_items(entry, url_map) + items = build_media_items(entry, url_map) data["metrics"] = [{"index": current_step or 0, "items": items}] results.append(data) return results def _fetch_media_all(self) -> List[Dict[str, Any]]: """获取全部媒体数据,后端 columns 批量一次返回。""" - payload = Metric._build_media_payload( + payload = build_media_payload( self._project_id, self._run_id, self._created_at, @@ -1089,7 +883,7 @@ def _fetch_media_all(self) -> List[Dict[str, Any]]: if isinstance(m, dict) for p in (m.get("data") or []) ] - url_map = Metric._fetch_presigned_urls(self, prefix, all_paths) if all_paths else {} + url_map = fetch_presigned_urls(self, prefix, all_paths) if all_paths else {} if all_paths: console.debug( f"Media fetched (all): run_id[{self._run_id}] - {len(all_paths)} items across {len(self._keys)} keys, requesting presigned urls..." @@ -1104,13 +898,17 @@ def _fetch_media_all(self) -> List[Dict[str, Any]]: "key": key, "metrics": [], } + if self._root_pro_id: + data["rootProId"] = self._root_pro_id + if self._root_exp_id: + data["rootExpId"] = self._root_exp_id entry = key_to_entry.get(key) if entry: metrics_list: List[Dict[str, Any]] = [] for m in entry.get("metrics") or []: if not isinstance(m, dict): continue - items = Metric._build_media_items(m, url_map) + items = build_media_items(m, url_map) metrics_list.append({"index": m.get("index", 0), "items": items}) data["metrics"] = metrics_list results.append(data) diff --git a/swanlab/api/project.py b/swanlab/api/project.py index 745003923..1968fa369 100644 --- a/swanlab/api/project.py +++ b/swanlab/api/project.py @@ -8,9 +8,9 @@ from typing import Any, Dict, Iterator, List, Optional, cast from swanlab.api.base import ApiClientContext, BaseEntity +from swanlab.api.helper import get_properties from swanlab.api.typings.common import PaginatedQuery from swanlab.api.typings.project import ApiProjectCountType, ApiProjectLabelType, ApiProjectType -from swanlab.api.utils import get_properties from swanlab.sdk.internal.pkg import console diff --git a/swanlab/api/self_hosted.py b/swanlab/api/self_hosted.py index f662fc9eb..4107de330 100644 --- a/swanlab/api/self_hosted.py +++ b/swanlab/api/self_hosted.py @@ -8,9 +8,9 @@ from typing import Any, Dict, Iterator, Optional, cast from swanlab.api.base import ApiClientContext, BaseEntity +from swanlab.api.helper import get_properties, validate_non_empty_string from swanlab.api.typings.common import ApiResponseType, PaginatedQuery from swanlab.api.typings.selfhosted import ApiLicensePlanLiteral, ApiSelfHostedInfoType -from swanlab.api.utils import get_properties, validate_non_empty_string class SelfHosted(BaseEntity): diff --git a/swanlab/api/series.py b/swanlab/api/series.py index 5a82775a4..65d1360c5 100644 --- a/swanlab/api/series.py +++ b/swanlab/api/series.py @@ -8,8 +8,8 @@ from typing import Any, Callable, Dict, Iterator, List, Optional from swanlab.api.base import ApiClientContext, BaseEntity +from swanlab.api.helper import fetch_file_presigned_urls, get_properties from swanlab.api.typings import ApiMetricKeyClassLiteral, ApiMetricKeyTypeLiteral, ApiResponseType -from swanlab.api.utils import get_properties # 系统指标 key 前缀:SCALAR 类型且以此前缀开头的 key 分类为 SYSTEM _SYSTEM_KEY_PREFIX = "__swanlab__" @@ -88,6 +88,7 @@ def metric( ignore_timestamp: bool = False, media_step: Optional[int] = None, all: bool = False, + x_axis: str = "step", ) -> Dict[str, Any]: """Fetch metric data points for this key. @@ -95,6 +96,9 @@ def metric( :param ignore_timestamp: If True, omit ``timestamp`` field from each data point. :param media_step: Step filter for MEDIA metrics. If None, returns all steps. :param all: If True, fetch full-resolution data without downsampling. + :param x_axis: X axis of the ``index`` values — ``"step"`` (default), a built-in axis + (``"time"`` / ``"relative_time"``), or any other non-empty string as a custom + x column key (SCALAR only). :returns: ``{"list": [{"step", "value", "timestamp", "key"}], ...}`` """ from swanlab.api.metric import Metric @@ -112,6 +116,7 @@ def metric( root_pro_id=self._root_pro_id, root_exp_id=self._root_exp_id, created_at=self._created_at, + x_axis=x_axis, ) return metric.json() @@ -121,7 +126,6 @@ def export_csv(self) -> ApiResponseType: :returns: ``ApiResponseType(ok=True, data={"url": ""})``. Returns ``ok=False`` for MEDIA keys. """ - from swanlab.api.metric import Metric if self._metric_type != "SCALAR": return ApiResponseType(ok=False, errmsg="export_csv() only support SCALAR metric_type", data=None) @@ -152,7 +156,7 @@ def export_csv(self) -> ApiResponseType: if not cos_key: return ApiResponseType(ok=False, errmsg="Invalid response format: missing cosKey", data=None) - url_map = Metric._fetch_file_presigned_urls(self, [cos_key]) + url_map = fetch_file_presigned_urls(self, [cos_key]) url = url_map.get(cos_key, "") if not url: return ApiResponseType(ok=False, errmsg="Failed to get presigned download URL", data=None) diff --git a/swanlab/api/typings/__init__.py b/swanlab/api/typings/__init__.py index 9764cb21b..77c11e9c9 100644 --- a/swanlab/api/typings/__init__.py +++ b/swanlab/api/typings/__init__.py @@ -26,6 +26,9 @@ from .metric import ( ApiLogSeriesType, ApiMediaSeriesType, + ApiMetricXAxisKindLiteral, + ApiMetricXAxisParam, + ApiMetricXAxisType, ApiScalarSeriesType, ) from .project import ApiProjectCountType, ApiProjectLabelType, ApiProjectType @@ -73,4 +76,7 @@ "ApiLogSeriesType", "ApiMediaSeriesType", "ApiScalarSeriesType", + "ApiMetricXAxisKindLiteral", + "ApiMetricXAxisParam", + "ApiMetricXAxisType", ] diff --git a/swanlab/api/typings/common.py b/swanlab/api/typings/common.py index b00a0eab3..947005157 100644 --- a/swanlab/api/typings/common.py +++ b/swanlab/api/typings/common.py @@ -78,7 +78,7 @@ # 指标日志级别 ApiMetricLogLevelLiteral = Literal["DEBUG", "INFO", "WARN", "ERROR"] -# X 轴类型 +# X 轴内置类型 ApiMetricXAxisLiteral = Literal["step", "time", "relative_time"] @@ -154,7 +154,9 @@ class RangeQuery(BaseModel, frozen=True): """ Scalar metric range query parameters. - type: Filter axis — ``"step"`` (default) or ``"timestamp"`` + type: Filter axis — ``"step"`` (default), ``"timestamp"``, or ``"custom"`` + (filters on the custom x-axis value domain; only valid when the queried + keys resolve to a custom x axis) start: Lower bound (inclusive); None = from beginning end: Upper bound (inclusive); None = to end last: Last N milliseconds (mutually exclusive with start/end; SDK auto-converts to timestamp filter) @@ -162,17 +164,33 @@ class RangeQuery(BaseModel, frozen=True): tail: Last N data points (mutually exclusive with head) head/tail can be combined with start/end/last: range filter first, then head/tail. + + ``start`` / ``end`` are stored as ``float`` uniformly (int inputs coerced by + pydantic; downstream comparisons are float-vs-float). Numeric rules depend on + ``type``: ``step`` / ``timestamp`` require non-negative integral values; + ``custom`` allows any finite float, including negative ones. NaN / ±Inf bounds + are rejected in every mode (a NaN bound would silently disable its filter). """ - type: Literal["step", "timestamp"] = "step" - start: Optional[int] = Field(default=None, ge=0) - end: Optional[int] = Field(default=None, ge=0) + type: Literal["step", "timestamp", "custom"] = "step" + start: Optional[float] = Field(default=None, allow_inf_nan=False) + end: Optional[float] = Field(default=None, allow_inf_nan=False) last: Optional[int] = Field(default=None, gt=0) head: Optional[int] = Field(default=None, gt=0) tail: Optional[int] = Field(default=None, gt=0) @model_validator(mode="after") def _validate_range_query(self) -> "RangeQuery": + if self.type in ("step", "timestamp"): + for name in ("start", "end"): + value = getattr(self, name) + if value is None: + continue + if value % 1 != 0 or value < 0: + raise ValueError( + f"{name} must be a non-negative integer for type {self.type!r}, got {value!r}; " + "fractional or negative bounds are only allowed with type='custom'" + ) if self.head is not None and self.tail is not None: raise ValueError("head and tail are mutually exclusive") if self.start is not None and self.end is not None and self.start > self.end: diff --git a/swanlab/api/typings/metric.py b/swanlab/api/typings/metric.py index e4a6fb3b8..cda5fd9cf 100644 --- a/swanlab/api/typings/metric.py +++ b/swanlab/api/typings/metric.py @@ -5,7 +5,9 @@ @description: 指标数据类型定义(用于 column 采样值) """ -from typing import Any, List, TypedDict, Union +from typing import Any, List, Literal, TypedDict, Union + +from swanlab.api.typings.common import ApiMetricXAxisLiteral # --------------------------------------------------------------------------- # Common — 通用指标类型定义 @@ -16,6 +18,23 @@ ApiMetricValueType = Union[int, float, str] +# --------------------------------------------------------------------------- +# X Axis — 查询参数与回显标记 +# --------------------------------------------------------------------------- +# x_axis 查询参数:内置轴字面量,或自定义 x 列 key(任意非空字符串) +ApiMetricXAxisParam = Union[ApiMetricXAxisLiteral, str] + +# 回显中 axis 的类别:"step" 为内置步数轴;CUSTOM/SYSTEM 携带 key +ApiMetricXAxisKindLiteral = Literal["step", "CUSTOM", "SYSTEM"] + + +class ApiMetricXAxisType(TypedDict, total=False): + """描述 ``metrics[].index`` 的实际语义,随每个 per-key metric 回显。""" + + type: ApiMetricXAxisKindLiteral + key: str + + # --------------------------------------------------------------------------- # Column Reference — 指标列引用,标识要查询的指标列 # --------------------------------------------------------------------------- @@ -31,10 +50,14 @@ class ApiMetricColumnRefType(TypedDict, total=False): # Scalar — 标量指标类型 # --------------------------------------------------------------------------- # 使用 index 因为 x 轴可以是 step / time / relative_time / 自定义列 +# 采样响应使用 data,CSV 导出记录使用 value; +# step 仅在 CSV 全量路径(all / range_query)下填充 class ApiScalarType(TypedDict, total=False): index: float data: ApiMetricValueType + value: ApiMetricValueType timestamp: int + step: int # 组合 /metrics/scalar 和 /metrics/scalar/value 的标量序列 @@ -43,6 +66,7 @@ class ApiScalarSeriesType(ApiMetricColumnRefType, total=False): metrics: List[ApiScalarType] url: str + xAxis: ApiMetricXAxisType min: ApiScalarType max: ApiScalarType avg: ApiScalarType diff --git a/swanlab/api/user.py b/swanlab/api/user.py index 67f6afd07..7f07218c1 100644 --- a/swanlab/api/user.py +++ b/swanlab/api/user.py @@ -8,8 +8,8 @@ from typing import Any, Dict, Optional from swanlab.api.base import ApiClientContext, BaseEntity +from swanlab.api.helper import get_properties, strip_dict from swanlab.api.typings.user import ApiUserProfileType -from swanlab.api.utils import get_properties, strip_dict class User(BaseEntity): diff --git a/swanlab/api/workspace.py b/swanlab/api/workspace.py index 0ef891d54..a9fb991b4 100644 --- a/swanlab/api/workspace.py +++ b/swanlab/api/workspace.py @@ -8,10 +8,10 @@ from typing import TYPE_CHECKING, Any, Dict, Iterator, List, Optional, cast from swanlab.api.base import ApiClientContext, BaseEntity +from swanlab.api.helper import get_properties, strip_dict, validate_project_name, validate_visibility from swanlab.api.typings.common import ApiVisibilityLiteral, PaginatedQuery from swanlab.api.typings.project import ApiProjectType from swanlab.api.typings.workspace import ApiWorkspaceLiteral, ApiWorkspaceProfileType, ApiWorkspaceType -from swanlab.api.utils import get_properties, strip_dict, validate_project_name, validate_visibility if TYPE_CHECKING: from swanlab.api.project import Project diff --git a/swanlab/cli/api/experiment.py b/swanlab/cli/api/experiment.py index ffa6e896c..659d43403 100644 --- a/swanlab/cli/api/experiment.py +++ b/swanlab/cli/api/experiment.py @@ -211,26 +211,39 @@ def list_experiment_series( help="Remove timestamp from metric data.", ) @click.option("--all", "fetch_all", is_flag=True, default=False, help="Fetch all data (CSV export for scalars).") +@click.option( + "--x-axis", + "x_axis", + default="step", + type=str, + help=("X axis of index values: 'step' (default), 'time' / 'relative_time', or a custom x column key."), +) @click.option( "--range-type", "range_type", default=None, - type=click.Choice(["step", "timestamp"], case_sensitive=False), - help="Range query type: 'step' or 'timestamp'.", + type=click.Choice(["step", "timestamp", "custom"], case_sensitive=False), + help="Range query type: 'step', 'timestamp', or 'custom' (custom x value domain; requires a custom x axis).", ) @click.option( "--range-start", "range_start", default=None, - type=click.IntRange(min=0), - help="Range start (inclusive). Step number or Unix timestamp in milliseconds.", + type=click.FLOAT, + help=( + "Range start (inclusive). Step number, Unix timestamp in milliseconds, or custom x value " + "(float allowed). Non-negative integer required for step/timestamp." + ), ) @click.option( "--range-end", "range_end", default=None, - type=click.IntRange(min=0), - help="Range end (inclusive). Step number or Unix timestamp in milliseconds.", + type=click.FLOAT, + help=( + "Range end (inclusive). Step number, Unix timestamp in milliseconds, or custom x value " + "(float allowed). Non-negative integer required for step/timestamp." + ), ) @click.option("--range-head", "range_head", default=None, type=click.IntRange(min=1), help="First N data points.") @click.option("--range-tail", "range_tail", default=None, type=click.IntRange(min=1), help="Last N data points.") @@ -248,9 +261,10 @@ def get_experiment_metrics( sample: int, ignore_timestamp: bool, fetch_all: bool, + x_axis: str, range_type: Optional[str], - range_start: Optional[int], - range_end: Optional[int], + range_start: Optional[float], + range_end: Optional[float], range_head: Optional[int], range_tail: Optional[int], range_last: Optional[int], @@ -294,6 +308,7 @@ def get_experiment_metrics( root_pro_id=experiment.root_pro_id, root_exp_id=experiment.root_exp_id, created_at=experiment.created_at_ts, + x_axis=x_axis, ).wrapper() diff --git a/tests/unit/api/test_api.py b/tests/unit/api/test_api.py index 313a811f2..ab40437fc 100644 --- a/tests/unit/api/test_api.py +++ b/tests/unit/api/test_api.py @@ -289,7 +289,9 @@ def get(path, **kwargs): exp.metrics(["loss"]) payload = ctx.client.post.call_args_list[0].kwargs["data"] - assert [call.args[0] for call in ctx.client.get.call_args_list] == ["/project/user/proj/runs/run-slug"] + assert [call.args[0] for call in ctx.client.get.call_args_list] == [ + "/project/user/proj/runs/run-slug", + ] assert payload["projectId"] == "project-cuid" assert payload["columns"] == [{"experimentId": "run-cuid", "key": "loss", "createdAt": 1722470400}] @@ -628,3 +630,111 @@ def test_invalid_sort_raises_on_iter(self, ctx): exps = Experiments(ctx, path="user/proj", sorts=bad_sorts, mode="post") with pytest.raises(ValueError, match="Invalid sort order"): list(exps) + + +# --------------------------------------------------------------------------- +# Metrics 自定义 X 轴单测 +# --------------------------------------------------------------------------- +class TestMetricsWithCustomXAxis: + def test_metrics_sampled_with_custom_x_axis(self, ctx): + def get(path, **kwargs): + if path == "/project/user/proj/runs/run-slug": + return _api_response( + { + "cuid": "run-cuid", + "slug": "run-slug", + "name": "test-run", + "project_id": "project-cuid", + "createdAt": "2024-08-01T00:00:00Z", + } + ) + raise AssertionError(f"unexpected GET {path}") + + ctx.client.get.side_effect = get + ctx.client.post.side_effect = [ + _api_response([{"metrics": [{"step": 1, "value": 0.5}], "key": "loss"}]), + _api_response([{"min": {"value": 0.5}}]), + ] + exp = Experiment(ctx, path="user/proj/run-slug") + res = exp.metrics(["loss"], x_axis="epoch") + + payload = ctx.client.post.call_args_list[0].kwargs["data"] + assert payload["xKey"] == "epoch" + assert payload["xType"] == "step" + + loss_item = res["list"][0] + assert loss_item["xAxis"] == {"type": "CUSTOM", "key": "epoch"} + + def test_metrics_with_x_axis_in_keys_identity(self, ctx): + """当 x_axis='epoch' 且 'epoch' 在 keys 中,epoch 序列的 value == index。""" + + def get(path, **kwargs): + return _api_response( + { + "cuid": "run-cuid", + "slug": "run-slug", + "name": "test-run", + "project_id": "project-cuid", + "createdAt": "2024-08-01T00:00:00Z", + } + ) + + ctx.client.get.side_effect = get + ctx.client.post.side_effect = [ + _api_response( + [ + {"metrics": [{"index": 10.0, "data": 10.0}], "key": "epoch"}, + {"metrics": [{"index": 10.0, "data": 0.5}], "key": "loss"}, + ] + ), + _api_response([{"min": {"value": 10.0}}, {"min": {"value": 0.5}}]), + ] + exp = Experiment(ctx, path="user/proj/run-slug") + res = exp.metrics(["epoch", "loss"], x_axis="epoch") + + assert len(res["list"]) == 2 + epoch_series = res["list"][0] + loss_series = res["list"][1] + + assert epoch_series["key"] == "epoch" + for pt in epoch_series["metrics"]: + assert pt["data"] == pt["index"] + + assert loss_series["key"] == "loss" + assert loss_series["xAxis"] == {"type": "CUSTOM", "key": "epoch"} + + def test_metrics_rejects_time_axis_in_csv_mode(self, ctx): + ctx.client.get.return_value = _api_response( + {"cuid": "run-cuid", "project_id": "p-1", "createdAt": "2024-08-01T00:00:00Z"} + ) + exp = Experiment(ctx, path="user/proj/run-slug") + with pytest.raises(ValueError, match="is not supported in CSV mode"): + exp.metrics(["loss"], all=True, x_axis="time") + + def test_metrics_rejects_custom_range_without_custom_x_axis(self, ctx): + ctx.client.get.return_value = _api_response( + {"cuid": "run-cuid", "project_id": "p-1", "createdAt": "2024-08-01T00:00:00Z"} + ) + exp = Experiment(ctx, path="user/proj/run-slug") + with pytest.raises(ValueError, match="requires a custom x axis"): + exp.metrics(["loss"], range_query={"type": "custom", "start": 0}, x_axis="step") + + def test_metrics_retains_root_ids(self, ctx): + ctx.client.post.side_effect = [ + _api_response([{"metrics": [{"step": 1, "value": 0.5}], "key": "loss"}]), + _api_response([{"min": {"value": 0.5}}]), + ] + metrics = Metrics( + ctx, + project_id="p-1", + run_id="r-1", + keys=["loss"], + metric_type="SCALAR", + root_pro_id="root-p", + root_exp_id="root-r", + created_at=1700000000, + ) + res = list(metrics) + assert len(res) == 1 + assert res[0].json()["rootProId"] == "root-p" + assert res[0].json()["rootExpId"] == "root-r" diff --git a/tests/unit/api/test_extractor.py b/tests/unit/api/test_extractor.py new file mode 100644 index 000000000..ad10acaf6 --- /dev/null +++ b/tests/unit/api/test_extractor.py @@ -0,0 +1,167 @@ +""" +@author: caddiesnew +@time: 2026/9/3 +@description: swanlab/api/helper/extractor.py 核心流式 CSV 解析与对齐单测 +""" + +import math +from unittest.mock import MagicMock + +from swanlab.api.helper.extractor import stream_export_csv +from swanlab.api.typings.common import RangeQuery + + +def _mock_client(csv_content: str): + mock_resp = MagicMock() + mock_resp.raise_for_status.return_value = None + mock_resp.encoding = "utf-8" + mock_resp.iter_lines.return_value = iter(csv_content.strip().splitlines()) + + mock_client = MagicMock() + mock_client._session.get.return_value = mock_resp + return mock_client + + +class TestStreamExportCsv: + def test_non_custom_x_axis_skips_missing_cells(self): + """非自定义轴(纯 step 轴,not x_key): + 旧契约守恒,空单元格或无效值直接跳过,不填充 NaN。 + """ + csv_text = ( + "step,c1,c1_timestamp,c2,c2_timestamp\n" + "0,0.5,1700000000000,0.8,1700000000000\n" + "1,0.4,1700000001000,,1700000001000\n" + "2,,1700000002000,0.9,1700000002000\n" + ) + client = _mock_client(csv_text) + result = stream_export_csv( + client=client, + url="https://mock/csv", + keys=["loss", "acc"], + rq=None, + x_key="", + ) + assert result is not None + loss_metrics = result["loss"] + acc_metrics = result["acc"] + + assert len(loss_metrics) == 2 + assert [m["step"] for m in loss_metrics] == [0, 1] + assert [m["value"] for m in loss_metrics] == [0.5, 0.4] + assert not any(math.isnan(m["value"]) for m in loss_metrics) + + assert len(acc_metrics) == 2 + assert [m["step"] for m in acc_metrics] == [0, 2] + assert [m["value"] for m in acc_metrics] == [0.8, 0.9] + assert not any(math.isnan(m["value"]) for m in acc_metrics) + + def test_custom_x_axis_nan_placeholders_and_index(self): + """自定义 x 轴(x_key 存在且不在 keys 中): + x_key 作为追加列位于末尾。缺失单元格以 NaN 占位,index 为对应的 x_value。 + """ + csv_text = ( + "step,c1,c1_ts,c2,c2_ts,c3,c3_ts\n" + "0,0.5,1000,0.8,1000,1.0,1000\n" + "1,0.4,2000,,2000,2.0,2000\n" + "2,,3000,0.9,3000,3.0,3000\n" + ) + client = _mock_client(csv_text) + result = stream_export_csv( + client=client, + url="https://mock/csv", + keys=["loss", "acc"], + rq=None, + x_key="epoch", + ) + assert result is not None + loss_metrics = result["loss"] + acc_metrics = result["acc"] + + assert len(loss_metrics) == 3 + assert len(acc_metrics) == 3 + + assert [m["index"] for m in loss_metrics] == [1.0, 2.0, 3.0] + assert [m["index"] for m in acc_metrics] == [1.0, 2.0, 3.0] + + assert math.isnan(acc_metrics[1]["value"]) + assert math.isnan(loss_metrics[2]["value"]) + + def test_x_key_in_keys_column_reuse_and_identity(self): + """当 x_key 存在于 keys 中(如 keys=["epoch", "loss"], x_key="epoch"): + 列复用,不追加末尾列,且 epoch 序列的 value == index。 + """ + csv_text = "step,c1,c1_ts,c2,c2_ts\n0,10.0,1000,0.5,1000\n1,20.0,2000,0.4,2000\n" + client = _mock_client(csv_text) + result = stream_export_csv( + client=client, + url="https://mock/csv", + keys=["epoch", "loss"], + rq=None, + x_key="epoch", + ) + assert result is not None + epoch_metrics = result["epoch"] + loss_metrics = result["loss"] + + for m in epoch_metrics: + assert m["value"] == m["index"] + assert [m["index"] for m in epoch_metrics] == [10.0, 20.0] + + assert [m["index"] for m in loss_metrics] == [10.0, 20.0] + assert [m["value"] for m in loss_metrics] == [0.5, 0.4] + + def test_custom_range_bounds_and_nan_drop(self): + """type="custom" 时: + 按 x_value 过滤 start/end,支持负值范围; + 当 x 缺失为 NaN/None 时,在有界查询下被过滤丢弃。 + """ + csv_text = ( + "step,c1,c1_ts,c2,c2_ts\n0,0.1,1000,-1.0,1000\n1,0.2,2000,,2000\n2,0.3,3000,0.5,3000\n3,0.4,4000,2.0,4000\n" + ) + client = _mock_client(csv_text) + rq = RangeQuery(type="custom", start=-0.5, end=1.0) + result = stream_export_csv( + client=client, + url="https://mock/csv", + keys=["loss"], + rq=rq, + x_key="lr", + ) + assert result is not None + metrics = result["loss"] + assert len(metrics) == 1 + assert metrics[0]["index"] == 0.5 + assert metrics[0]["value"] == 0.3 + + def test_range_head_and_tail_with_custom_x(self): + """head 和 tail 在 custom x 轴下的截断行为。""" + csv_text = ( + "step,c1,c1_ts,c2,c2_ts\n" + "0,0.1,1000,1.0,1000\n" + "1,0.2,2000,2.0,2000\n" + "2,0.3,3000,3.0,3000\n" + "3,0.4,4000,4.0,4000\n" + ) + # 测试 head=2 + client = _mock_client(csv_text) + result_head = stream_export_csv( + client=client, + url="https://mock/csv", + keys=["loss"], + rq=RangeQuery(type="custom", head=2), + x_key="epoch", + ) + assert result_head is not None + assert [m["index"] for m in result_head["loss"]] == [1.0, 2.0] + + # 测试 tail=2 + client = _mock_client(csv_text) + result_tail = stream_export_csv( + client=client, + url="https://mock/csv", + keys=["loss"], + rq=RangeQuery(type="custom", tail=2), + x_key="epoch", + ) + assert result_tail is not None + assert [m["index"] for m in result_tail["loss"]] == [3.0, 4.0] diff --git a/tests/unit/api/test_utils.py b/tests/unit/api/test_utils.py index e93d2c75b..b895c2291 100644 --- a/tests/unit/api/test_utils.py +++ b/tests/unit/api/test_utils.py @@ -8,10 +8,11 @@ import pytest -from swanlab.api.self_hosted import SelfHosted -from swanlab.api.typings.common import PaginatedQuery -from swanlab.api.typings.selfhosted import ApiSelfHostedInfoType -from swanlab.api.utils import ( +from swanlab.api.helper import ( + RELATIVE_TIME_AXIS, + STEP_AXIS, + TIME_AXIS, + builtin_x_axis, parse_timestamp_ms, validate_column_params, validate_filter, @@ -20,7 +21,11 @@ validate_metric_type, validate_project_name, validate_sort, + validate_x_axis, ) +from swanlab.api.self_hosted import SelfHosted +from swanlab.api.typings.common import PaginatedQuery, RangeQuery +from swanlab.api.typings.selfhosted import ApiSelfHostedInfoType # --------------------------------------------------------------------------- @@ -221,3 +226,104 @@ def test_unsupported_type_raises(self): def test_float_raises(self): with pytest.raises(ValueError, match="Expected str or int"): parse_timestamp_ms(1715769600.0) # type: ignore[arg-type] + + +# --------------------------------------------------------------------------- +# validate_x_axis +# --------------------------------------------------------------------------- +class TestValidateXAxis: + @pytest.mark.parametrize("axis", ["step", "time", "relative_time"]) + def test_valid_builtins(self, axis: str): + assert validate_x_axis(axis) == axis + + @pytest.mark.parametrize("axis", ["epoch", "iter", "custom_axis", "123", "lr/step"]) + def test_valid_custom_key(self, axis: str): + assert validate_x_axis(axis) == axis + + @pytest.mark.parametrize("axis", ["", " ", "\t\n"]) + def test_reject_empty_or_whitespace(self, axis: str): + with pytest.raises(ValueError, match="non-empty string"): + validate_x_axis(axis) + + @pytest.mark.parametrize("axis", [None, 123, ["epoch"]]) + def test_reject_non_string(self, axis): + with pytest.raises(ValueError, match="non-empty string"): + validate_x_axis(axis) # type: ignore[arg-type] + + @pytest.mark.parametrize("metric_type", ["MEDIA", "LOG"]) + def test_non_scalar_accepts_step(self, metric_type: str): + assert validate_x_axis("step", metric_type=metric_type) == "step" + + @pytest.mark.parametrize("metric_type", ["MEDIA", "LOG"]) + @pytest.mark.parametrize("axis", ["time", "relative_time"]) + def test_non_scalar_rejects_builtin_time_axes(self, metric_type: str, axis: str): + with pytest.raises(ValueError, match="only supported for SCALAR metrics"): + validate_x_axis(axis, metric_type=metric_type) + + @pytest.mark.parametrize("metric_type", ["MEDIA", "LOG"]) + def test_non_scalar_rejects_custom_key(self, metric_type: str): + with pytest.raises(ValueError, match="custom x_axis key 'epoch' is only supported for SCALAR metrics"): + validate_x_axis("epoch", metric_type=metric_type) + + +# --------------------------------------------------------------------------- +# builtin_x_axis & axis_request_params +# --------------------------------------------------------------------------- +class TestBuiltinXAxis: + def test_builtins(self): + assert builtin_x_axis("step") == STEP_AXIS + assert builtin_x_axis("time") == TIME_AXIS + assert builtin_x_axis("relative_time") == RELATIVE_TIME_AXIS + + @pytest.mark.parametrize("axis", ["epoch", "auto", "timestamp", "unknown", ""]) + def test_non_builtins_return_none(self, axis: str): + assert builtin_x_axis(axis) is None + + +# --------------------------------------------------------------------------- +# RangeQuery custom & bounds validation +# --------------------------------------------------------------------------- +class TestRangeQueryBounds: + def test_step_integer_bounds_accepted(self): + rq = RangeQuery(type="step", start=0, end=100) + assert rq.start == 0.0 + assert rq.end == 100.0 + + rq_float_int = RangeQuery(type="step", start=5.0, end=10.0) + assert rq_float_int.start == 5.0 + assert rq_float_int.end == 10.0 + + @pytest.mark.parametrize("invalid_val", [1.5, -1, -0.1]) + def test_step_rejects_fraction_or_negative(self, invalid_val): + with pytest.raises(ValueError, match="must be a non-negative integer for type 'step'"): + RangeQuery(type="step", start=invalid_val) + with pytest.raises(ValueError, match="must be a non-negative integer for type 'step'"): + RangeQuery(type="step", end=invalid_val) + + def test_timestamp_bounds(self): + rq = RangeQuery(type="timestamp", start=1700000000000, end=1700000001000) + assert rq.start == 1700000000000.0 + + with pytest.raises(ValueError, match="must be a non-negative integer for type 'timestamp'"): + RangeQuery(type="timestamp", start=-100) + with pytest.raises(ValueError, match="must be a non-negative integer for type 'timestamp'"): + RangeQuery(type="timestamp", end=10.5) + + def test_custom_bounds_allow_negative_and_floats(self): + rq = RangeQuery(type="custom", start=-10.5, end=1.25e-3) + assert rq.start == -10.5 + assert rq.end == 1.25e-3 + + def test_custom_bounds_start_greater_than_end_raises(self): + with pytest.raises(ValueError, match="start must be <= end"): + RangeQuery(type="custom", start=10.0, end=5.0) + + def test_head_and_tail_mutually_exclusive(self): + with pytest.raises(ValueError, match="head and tail are mutually exclusive"): + RangeQuery(type="step", head=10, tail=10) + + def test_last_and_start_end_mutually_exclusive(self): + with pytest.raises(ValueError, match="last is mutually exclusive with start/end"): + RangeQuery(type="step", last=1000, start=10) + with pytest.raises(ValueError, match="last is mutually exclusive with start/end"): + RangeQuery(type="step", last=1000, end=20) From 4af0bf9db8071ca319b4dd3e0ea065ec7844fb4f Mon Sep 17 00:00:00 2001 From: CaddiesNew <50736785+Nexisato@users.noreply.github.com> Date: Sun, 6 Sep 2026 18:54:17 +0800 Subject: [PATCH 09/11] test(benchmark): add async log bench (#1769) * test(benchmark): add async log bench * test(benchmark): fix pyright possibly-unbound errors in async log bench --------- Co-authored-by: Kang Li <79990647+SAKURA-CAT@users.noreply.github.com> --- tests/benchmark/sdk/cmd/bench_async_log.py | 473 +++++++++++++++++++++ 1 file changed, 473 insertions(+) create mode 100644 tests/benchmark/sdk/cmd/bench_async_log.py diff --git a/tests/benchmark/sdk/cmd/bench_async_log.py b/tests/benchmark/sdk/cmd/bench_async_log.py new file mode 100644 index 000000000..5acb4feb7 --- /dev/null +++ b/tests/benchmark/sdk/cmd/bench_async_log.py @@ -0,0 +1,473 @@ +""" +@author: cunyue, nexisato +@file: bench_async_log.py +@time: 2026/9/4 +@description: swanlab.log vs swanlab.async_log (threading 模式) 性能对比基准测试 + +对比维度: + 1. 上报模式: + - no_sdk: 纯计算基线(无任何 SwanLab 开销,仅 macro 作为 0% 黄金基准) + - sync_log: swanlab.log(data, step=step) 同步上报 + - async_log: swanlab.async_log(lambda: data, step=step, mode="threading") 异步上报 + 2. 指标规模: + - 10 keys: 常规训练常用标量(loss, acc, lr, 层权重/梯度模长等) + - 100 keys: 宽/层级命名空间标量(大模型各层细节监控) + 3. 评测类型: + - micro: 纯主线程调用耗时(Mean / P50 / P95 / P99 延迟,Throughput QPS) + - macro: 模拟训练循环(Step 耗时分布,相对 No-SDK 的 Overhead % 开销占比) + +运行方式: + # 1. 作为 pytest 基准测试运行: + uv run pytest tests/benchmark/sdk/cmd/bench_async_log.py -v -s + + # 2. 作为独立脚本运行(支持完整 CLI 参数): + uv run python tests/benchmark/sdk/cmd/bench_async_log.py + uv run python tests/benchmark/sdk/cmd/bench_async_log.py --steps 500 --workload-ms 10 +""" + +from __future__ import annotations + +import argparse +import json +import os +import shutil +import subprocess +import sys +import tempfile +import time +from typing import Any, Callable, Dict, List, Optional + +RESULT_MARK = "##BENCH_JSON## " + + +# =========================================================================== +# 指标数据生成器 (10 keys vs 100 keys) +# =========================================================================== + + +def generate_payload(num_keys: int, step: int) -> Dict[str, float]: + """生成指定 key 数量的指标字典,包含层级路径。""" + if num_keys == 10: + return { + "train/loss": 0.5 + 0.001 * (step % 100), + "train/accuracy": 0.85 + 0.0005 * (step % 100), + "train/learning_rate": 1e-4, + "train/grad_norm": 1.25, + "train/epoch": step / 100.0, + "layer1/weight_norm": 0.35, + "layer1/bias_norm": 0.05, + "layer2/weight_norm": 0.45, + "layer2/bias_norm": 0.06, + "val/loss": 0.52 + 0.001 * (step % 100), + } + elif num_keys == 100: + data: Dict[str, float] = {} + data["train/loss"] = 0.5 + 0.001 * (step % 100) + data["train/accuracy"] = 0.85 + 0.0005 * (step % 100) + data["train/lr"] = 1e-4 + data["train/epoch"] = step / 100.0 + # 其余 96 个指标按 12 个模块 × 8 个统计量组织 + stats = ["w_mean", "w_std", "b_mean", "b_std", "grad_norm", "grad_std", "act_mean", "act_std"] + for module_idx in range(12): + for stat in stats: + data[f"layers/block_{module_idx:02d}/{stat}"] = 0.01 * (step % 50) + module_idx * 0.1 + return data + else: + return {f"metric_{i:03d}": 1.0 + i * 0.01 for i in range(num_keys)} + + +# =========================================================================== +# 统计辅助函数 +# =========================================================================== + + +def percentile(sorted_data: List[float], pct: float) -> float: + """计算百分位数,pct ∈ [0, 100]""" + if not sorted_data: + return 0.0 + idx = int(len(sorted_data) * pct / 100.0) + return sorted_data[min(idx, len(sorted_data) - 1)] + + +def simulate_training_workload(workload_ms: float) -> None: + """模拟真实训练计算(释放 GIL 的等待 + 少量 CPU 计算)。""" + if workload_ms <= 0: + return + # 模拟真实混合计算:80% sleep (模拟 GPU forward/backward C 扩展), 20% CPU 占用 + sleep_s = (workload_ms * 0.8) / 1000.0 + time.sleep(sleep_s) + cpu_deadline = time.perf_counter() + (workload_ms * 0.2) / 1000.0 + val = 1.0001 + while time.perf_counter() < cpu_deadline: + val = (val * 1.000001) % 1000.0 + + +# =========================================================================== +# 子进程 Worker 实现(彻底隔离环境,消除全局单例和线程残留污染) +# =========================================================================== + + +def run_worker( + mode: str, + num_keys: int, + test_type: str, + steps: int, + warmup: int, + workload_ms: float, +) -> None: + """在独立子进程中执行单次压测。 + + :param mode: "no_sdk" | "sync_log" | "async_log" + :param num_keys: 10 | 100 + :param test_type: "micro" | "macro" + :param steps: 测量步骤数 + :param warmup: 预热步骤数(耗时不计入统计) + :param workload_ms: macro 模式下的每步模拟训练耗时 (ms) + """ + tmp_log_dir = tempfile.mkdtemp(prefix="swanlab_bench_") + call_latencies_ns: List[int] = [] + step_latencies_ns: List[int] = [] + + # 上报调用闭包,在 init 分支内绑定到具体实现,热循环内不做模式分支 + log_fn: Optional[Callable[[Dict[str, float], int], Any]] = None + finish_fn: Optional[Callable[[], Any]] = None + + try: + if mode != "no_sdk": + import swanlab + + # 统一使用 local 模式排除网络波动,设置 log_level 减少控制台干扰 + swanlab.init( + project="benchmarks", + experiment_name=f"bench_{mode}_{num_keys}k_{test_type}", + mode="local", + log_level="warning", + log_dir=tmp_log_dir, + ) + finish_fn = swanlab.finish + if mode == "sync_log": + + def sync_log_fn(d: Dict[str, float], s: int) -> None: + swanlab.log(d, step=s) + + log_fn = sync_log_fn + else: + + def async_log_fn(d: Dict[str, float], s: int) -> Any: + return swanlab.async_log(lambda d=d: d, step=s, mode="threading") + + log_fn = async_log_fn + + # ---------------------------------- + # 预热阶段 (Warmup) + # ---------------------------------- + for w in range(warmup): + data = generate_payload(num_keys, w) + if test_type == "macro": + simulate_training_workload(workload_ms) + if log_fn is not None: + log_fn(data, w) + + # ---------------------------------- + # 测量阶段 (Measurement) + # ---------------------------------- + t_total_start = time.perf_counter_ns() + + for s in range(warmup, warmup + steps): + data = generate_payload(num_keys, s) + + t_step_start = time.perf_counter_ns() + + if test_type == "macro": + simulate_training_workload(workload_ms) + + t_call_start = time.perf_counter_ns() + if log_fn is not None: + log_fn(data, s) + t_call_end = time.perf_counter_ns() + + t_step_end = time.perf_counter_ns() + + call_latencies_ns.append(t_call_end - t_call_start) + step_latencies_ns.append(t_step_end - t_step_start) + + # 主循环耗时到此为止;finish 的队列排空耗时单独统计,不计入吞吐量 + t_total_end = time.perf_counter_ns() + + t_finish_start = time.perf_counter_ns() + if finish_fn is not None: + finish_fn() + t_finish_end = time.perf_counter_ns() + + # ---------------------------------- + # 统计分析 + # ---------------------------------- + call_latencies_us = [v / 1e3 for v in call_latencies_ns] + step_latencies_ms = [v / 1e6 for v in step_latencies_ns] + + call_latencies_us_sorted = sorted(call_latencies_us) + step_latencies_ms_sorted = sorted(step_latencies_ms) + + total_time_s = (t_total_end - t_total_start) / 1e9 + finish_time_ms = (t_finish_end - t_finish_start) / 1e6 + + result = { + "mode": mode, + "num_keys": num_keys, + "test_type": test_type, + "steps": steps, + "warmup": warmup, + "workload_ms": workload_ms, + "total_time_s": round(total_time_s, 4), + "finish_time_ms": round(finish_time_ms, 2), + "throughput_qps": round(steps / total_time_s, 1), + # 单次调用延迟 (μs) + "call_mean_us": round(sum(call_latencies_us) / len(call_latencies_us), 2), + "call_p50_us": round(percentile(call_latencies_us_sorted, 50), 2), + "call_p95_us": round(percentile(call_latencies_us_sorted, 95), 2), + "call_p99_us": round(percentile(call_latencies_us_sorted, 99), 2), + # Step 耗时 (ms) + "step_mean_ms": round(sum(step_latencies_ms) / len(step_latencies_ms), 3), + "step_p50_ms": round(percentile(step_latencies_ms_sorted, 50), 3), + "step_p95_ms": round(percentile(step_latencies_ms_sorted, 95), 3), + "step_p99_ms": round(percentile(step_latencies_ms_sorted, 99), 3), + } + print(RESULT_MARK + json.dumps(result), flush=True) + + finally: + shutil.rmtree(tmp_log_dir, ignore_errors=True) + + +# =========================================================================== +# 父进程调用与结果汇总 +# =========================================================================== + + +def spawn_case( + mode: str, + num_keys: int, + test_type: str, + steps: int, + warmup: int, + workload_ms: float, +) -> Dict[str, Any]: + """在子进程中运行单个测试用例并解析结果。""" + cmd = [ + sys.executable, + os.path.abspath(__file__), + "--worker", + "--mode", + mode, + "--keys", + str(num_keys), + "--test-type", + test_type, + "--steps", + str(steps), + "--warmup", + str(warmup), + "--workload-ms", + str(workload_ms), + ] + proc = subprocess.run(cmd, capture_output=True, text=True, timeout=300) + payload: Optional[Dict[str, Any]] = None + for line in proc.stdout.splitlines(): + if line.startswith(RESULT_MARK): + payload = json.loads(line[len(RESULT_MARK) :]) + break + + if payload is None: + sys.stderr.write(proc.stderr) + raise RuntimeError(f"Worker failed: mode={mode}, keys={num_keys}, rc={proc.returncode}") + + return payload + + +def print_report( + micro_results: List[Dict[str, Any]], + macro_results: List[Dict[str, Any]], +) -> None: + """打印格式化的高亮性能对比报告。""" + print("\n" + "=" * 108) + print(" SwanLab Benchmark: swanlab.log vs swanlab.async_log (threading 模式)") + print("=" * 108) + + # 1. Micro Benchmark: 纯主线程单次调用延迟对比 + print("\n[Part 1] Micro-Benchmark: 纯主线程调用耗时与吞吐量 (无模拟训练开销)") + print("-" * 108) + print( + f" {'Keys':<8} {'Mode':<14} {'Call Mean(μs)':>14} {'P50(μs)':>10} {'P95(μs)':>10} {'P99(μs)':>10} {'Throughput(QPS)':>16} {'Finish(ms)':>12}" + ) + print("-" * 108) + + for r in micro_results: + print( + f" {r['num_keys']:<8} " + f"{r['mode']:<14} " + f"{r['call_mean_us']:>14.1f} " + f"{r['call_p50_us']:>10.1f} " + f"{r['call_p95_us']:>10.1f} " + f"{r['call_p99_us']:>10.1f} " + f"{r['throughput_qps']:>16.1f} " + f"{r['finish_time_ms']:>12.1f}" + ) + print("-" * 108) + + # 2. Macro Benchmark: 模拟训练循环耗时与开销占比对比 + print("\n[Part 2] Macro-Benchmark: 模拟训练 Step 耗时分布与开销占比 (%)") + print("-" * 108) + print( + f" {'Keys':<8} {'Mode':<14} {'Step Mean(ms)':>14} {'P95(ms)':>10} {'P99(ms)':>10} {'Overhead(%)':>13} {'vs sync_log':>16} {'Finish(ms)':>12}" + ) + print("-" * 108) + + # 按 num_keys 分组计算开销百分比 + grouped_keys: Dict[int, Dict[str, Dict[str, Any]]] = {} + for r in macro_results: + k = r["num_keys"] + grouped_keys.setdefault(k, {})[r["mode"]] = r + + for k, group in grouped_keys.items(): + base = group.get("no_sdk") + base_ms = base["step_mean_ms"] if base else 1.0 + sync_r = group.get("sync_log") + sync_overhead = (sync_r["step_mean_ms"] - base_ms) / base_ms * 100.0 if sync_r and base else 0.0 + + for mode_name in ("no_sdk", "sync_log", "async_log"): + if mode_name not in group: + continue + r = group[mode_name] + cur_ms = r["step_mean_ms"] + overhead_pct = (cur_ms - base_ms) / base_ms * 100.0 if base else 0.0 + + if mode_name == "no_sdk": + overhead_str = "0.00% (基准)" + vs_sync_str = "-" + elif mode_name == "sync_log": + overhead_str = f"+{overhead_pct:.2f}%" + vs_sync_str = "基准 (100%)" + else: + overhead_str = f"+{overhead_pct:.2f}%" + if sync_overhead > 0.001: + reduction = (sync_overhead - overhead_pct) / sync_overhead * 100.0 + vs_sync_str = f"开销减少 {reduction:.1f}%" + else: + vs_sync_str = "相当" + + print( + f" {k:<8} " + f"{mode_name:<14} " + f"{r['step_mean_ms']:>14.3f} " + f"{r['step_p95_ms']:>10.3f} " + f"{r['step_p99_ms']:>10.3f} " + f"{overhead_str:>13} " + f"{vs_sync_str:>16} " + f"{r['finish_time_ms']:>12.1f}" + ) + print("-" * 108) + + print("=" * 108 + "\n") + + +def run_full_benchmark( + micro_steps: int = 1000, + macro_steps: int = 300, + workload_ms: float = 10.0, + warmup: int = 50, +) -> Dict[str, Any]: + """执行完整的 10-key 和 100-key 对比测试矩阵。""" + keys_list = [10, 100] + modes_micro = ["sync_log", "async_log"] + modes_macro = ["no_sdk", "sync_log", "async_log"] + + micro_results: List[Dict[str, Any]] = [] + macro_results: List[Dict[str, Any]] = [] + + # 1. 运行 Micro-benchmark + for k in keys_list: + for m in modes_micro: + sys.stderr.write(f"[Bench] Micro test: keys={k}, mode={m}, steps={micro_steps}...\n") + res = spawn_case(m, k, "micro", micro_steps, warmup, workload_ms=0.0) + micro_results.append(res) + + # 2. 运行 Macro-benchmark + for k in keys_list: + for m in modes_macro: + sys.stderr.write( + f"[Bench] Macro test: keys={k}, mode={m}, steps={macro_steps}, workload={workload_ms}ms...\n" + ) + res = spawn_case(m, k, "macro", macro_steps, warmup, workload_ms=workload_ms) + macro_results.append(res) + + print_report(micro_results, macro_results) + return {"micro": micro_results, "macro": macro_results} + + +# =========================================================================== +# Pytest 测试入口 +# =========================================================================== + + +def test_bench_log_vs_async_log(): + """Pytest 基准测试套件集成入口。""" + # pytest 下适当缩短步数以保证测试快速通过,同时保留统计精度 + results = run_full_benchmark( + micro_steps=500, + macro_steps=200, + workload_ms=5.0, + warmup=30, + ) + # 验证测试完整性 + assert len(results["micro"]) == 4 # 2 keys * 2 modes + assert len(results["macro"]) == 6 # 2 keys * 3 modes + + # 验证 100-key 场景下 async_log 主线程单次调用延迟显著优于 sync_log + micro_100 = {r["mode"]: r for r in results["micro"] if r["num_keys"] == 100} + assert micro_100["async_log"]["call_mean_us"] < micro_100["sync_log"]["call_mean_us"], ( + "async_log should have lower main-thread latency than sync_log on 100 keys" + ) + + +# =========================================================================== +# CLI 独立运行入口 +# =========================================================================== + + +def main(): + parser = argparse.ArgumentParser(description="SwanLab: swanlab.log vs swanlab.async_log Benchmark") + parser.add_argument("--worker", action="store_true", help=argparse.SUPPRESS) + parser.add_argument("--mode", type=str, choices=["no_sdk", "sync_log", "async_log"], default="sync_log") + parser.add_argument("--keys", type=int, default=10, help="指标 key 数量 (10 或 100)") + parser.add_argument("--test-type", type=str, choices=["micro", "macro"], default="micro") + parser.add_argument("--steps", type=int, default=1000, help="测试 step 数") + parser.add_argument("--warmup", type=int, default=50, help="预热 step 数") + parser.add_argument("--workload-ms", type=float, default=10.0, help="模拟训练每步耗时 (ms)") + parser.add_argument("--out", type=str, default=None, help="导出结果 JSON 路径") + args = parser.parse_args() + + if args.worker: + run_worker( + mode=args.mode, + num_keys=args.keys, + test_type=args.test_type, + steps=args.steps, + warmup=args.warmup, + workload_ms=args.workload_ms, + ) + return + + results = run_full_benchmark( + micro_steps=args.steps, + macro_steps=max(200, args.steps // 3), + workload_ms=args.workload_ms, + warmup=args.warmup, + ) + + if args.out: + with open(args.out, "w", encoding="utf-8") as f: + json.dump(results, f, ensure_ascii=False, indent=2) + print(f"Results saved to {args.out}") + + +if __name__ == "__main__": + main() From 948849fe499ed849e1f6da644bd7940b86218a55 Mon Sep 17 00:00:00 2001 From: CaddiesNew <50736785+Nexisato@users.noreply.github.com> Date: Mon, 21 Sep 2026 22:30:16 +0800 Subject: [PATCH 10/11] fix: use named temp file for writability probe to avoid EIO on FUSE/NAS (#1774) --- swanlab/sdk/internal/pkg/fs/dir.py | 56 +++++++++-- tests/unit/sdk/cmd/init/test_init_e2e.py | 4 +- .../sdk/internal/pkg/fs/test_pkg_fs_dir.py | 97 +++++++++++++++++-- 3 files changed, 141 insertions(+), 16 deletions(-) diff --git a/swanlab/sdk/internal/pkg/fs/dir.py b/swanlab/sdk/internal/pkg/fs/dir.py index 52a0d85eb..f762273e5 100644 --- a/swanlab/sdk/internal/pkg/fs/dir.py +++ b/swanlab/sdk/internal/pkg/fs/dir.py @@ -11,7 +11,7 @@ from pathlib import Path from typing import Union -from .. import console +from swanlab.sdk.internal.pkg import console, safe def _get_fs_timeout(default: float = 5.0) -> float: @@ -36,6 +36,9 @@ def _get_fs_timeout(default: float = 5.0) -> float: # 模块加载时安全获取 TIMEOUT = _get_fs_timeout() +# 可写性探针临时文件前缀 +PROBE_PREFIX = ".swanlab_test_" + def safe_mkdirs(*paths: Union[str, Path], timeout: float = TIMEOUT, ensure_clean: bool = False): """ @@ -48,11 +51,47 @@ def safe_mkdirs(*paths: Union[str, Path], timeout: float = TIMEOUT, ensure_clean safe_mkdir(path, timeout=timeout, ensure_clean=ensure_clean) +def _probe_writable(p: Path) -> None: + """ + 目录可写性探针:创建命名文件 → 写入 → 关闭 → 删除。 + + 不使用 tempfile.TemporaryFile:其匿名文件(O_TMPFILE)或「fd 仍打开时立即 + unlink」的语义在部分 NAS 上会返回 EIO。只有创建、写入和关闭 + 用于判定目录可写;删除仅尽力清理本次探针文件。 + + 实现思路: + 1. 用 mkstemp 在目标目录创建命名探针文件。 + 2. 写入一个字节,验证文件可写。 + 3. 先关闭 fd,避免部分 FUSE / NAS 在文件仍打开时 unlink 返回 EIO。 + 4. 只尝试删除本次探针文件,不扫描或清理历史文件。 + 5. 创建、写入或关闭失败时上抛供 safe_mkdir 重试;删除失败仅记录 trace。 + """ + fd, name = tempfile.mkstemp(dir=p, prefix=PROBE_PREFIX) + try: + try: + os.write(fd, b"0") + except BaseException: + # close 的异常不得覆盖原始错误 + with safe.block(OSError, message=None): + os.close(fd) + raise + os.close(fd) + finally: + # 目录已经通过创建、写入和关闭完成可写性探测。 + # 清理失败不参与可写性判定,也不扫描或删除历史探针文件。 + with safe.block( + OSError, + message=f"Failed to clean up writability probe file [{name}]", + write_to_tty=False, + ): + os.unlink(name) + + def safe_mkdir(path: Union[str, Path], timeout: float = TIMEOUT, ensure_clean: bool = False) -> Path: """ 安全地创建目录,带有抗异步文件系统延迟的探针机制。 - 创建后会探测目录是否真正可见且可写,以容忍 NAS / NFS 等的异步IO延迟。 + 创建后会探测目录是否真正可见且可写,以容忍 NAS / NFS 等的异步IO延迟,可写性探测使用命名临时文件。 权限不足属于不可恢复错误,会立即抛出 PermissionError,不会重试。 :param path: 目录路径 @@ -84,20 +123,21 @@ def safe_mkdir(path: Union[str, Path], timeout: float = TIMEOUT, ensure_clean: b # 探测二:目录可见后可能仍暂不可写(权限问题立即失败,其余重试到超时) while True: try: - with tempfile.TemporaryFile(dir=p, prefix=".swanlab_test_") as f: - f.write(b"0") + _probe_writable(p) break - except PermissionError: - # 权限不足不会因重试而恢复,立即失败,避免被误判为可重试的文件系统延迟 + except PermissionError as e: + # 创建、写入或关闭探针文件被拒:权限不足不会因重试而恢复。 + console.trace(f"Directory [{p}] is not writable, underlying error: {e}") raise PermissionError( f"Directory [{p}] is not writable. Please choose a writable log_dir or update directory permissions." ) from None - except OSError: + except OSError as last_error: + # 取最后一次原始 OSError,避免把硬错误误报成单纯超时 if time.time() - start_time > timeout: console.trace(f"Directory {p} exists but is not writable within {timeout}s") raise TimeoutError( f"Directory [{p}] exists but is not writable within {timeout}s, FILESYSTEM may be slow or remote." - ) + ) from last_error time.sleep(0.1) return p diff --git a/tests/unit/sdk/cmd/init/test_init_e2e.py b/tests/unit/sdk/cmd/init/test_init_e2e.py index 09c649719..685a2c848 100644 --- a/tests/unit/sdk/cmd/init/test_init_e2e.py +++ b/tests/unit/sdk/cmd/init/test_init_e2e.py @@ -375,10 +375,10 @@ def test_init_fails_when_log_dir_is_not_writable(self, tmp_path, monkeypatch): log_dir = tmp_path / "root-owned-log-dir" log_dir.mkdir() - def deny_tempfile(*args, **kwargs): + def deny_mkstemp(*args, **kwargs): raise PermissionError(errno.EACCES, "Permission denied") - monkeypatch.setattr("swanlab.sdk.internal.pkg.fs.dir.tempfile.TemporaryFile", deny_tempfile) + monkeypatch.setattr("swanlab.sdk.internal.pkg.fs.dir.tempfile.mkstemp", deny_mkstemp) with pytest.raises(PermissionError, match="Directory .* is not writable"): init(mode="local", log_dir=str(log_dir)) diff --git a/tests/unit/sdk/internal/pkg/fs/test_pkg_fs_dir.py b/tests/unit/sdk/internal/pkg/fs/test_pkg_fs_dir.py index ac4e33c08..232a5f393 100644 --- a/tests/unit/sdk/internal/pkg/fs/test_pkg_fs_dir.py +++ b/tests/unit/sdk/internal/pkg/fs/test_pkg_fs_dir.py @@ -54,28 +54,113 @@ def test_safe_mkdir_nas_timeout(monkeypatch, tmp_path: Path): time_calls = [0, 10, 20, 30] monkeypatch.setattr("swanlab.sdk.internal.pkg.fs.dir.time.time", lambda: time_calls.pop(0)) - def mock_tempfile(*args, **kwargs): - raise OSError("Simulated NAS Permission Denied") + def mock_mkstemp(*args, **kwargs): + raise OSError("Simulated NAS IO Error") - monkeypatch.setattr("swanlab.sdk.internal.pkg.fs.dir.tempfile.TemporaryFile", mock_tempfile) + monkeypatch.setattr("swanlab.sdk.internal.pkg.fs.dir.tempfile.mkstemp", mock_mkstemp) - with pytest.raises(TimeoutError, match="is not writable within"): + with pytest.raises(TimeoutError, match="is not writable within") as exc_info: dir.safe_mkdir(target, timeout=5.0) + # 超时错误应链上最后一次原始 OSError,便于诊断根因 + assert isinstance(exc_info.value.__cause__, OSError) + def test_safe_mkdir_permission_denied_raises_permission_error(monkeypatch, tmp_path: Path): """权限不足时应立即抛 PermissionError,而不是伪装成 NAS 超时。""" target = tmp_path / "readonly_dir" - def mock_tempfile(*args, **kwargs): + def mock_mkstemp(*args, **kwargs): raise PermissionError(errno.EACCES, "Permission denied") - monkeypatch.setattr("swanlab.sdk.internal.pkg.fs.dir.tempfile.TemporaryFile", mock_tempfile) + monkeypatch.setattr("swanlab.sdk.internal.pkg.fs.dir.tempfile.mkstemp", mock_mkstemp) with pytest.raises(PermissionError, match="Directory .* is not writable"): dir.safe_mkdir(target, timeout=5.0) +def test_safe_mkdir_does_not_depend_on_temporary_file(monkeypatch, tmp_path: Path): + """依赖约束:可写性探测不得回退使用 tempfile.TemporaryFile。 + + 其匿名文件 / 「打开时 unlink」语义在部分 FUSE / NAS 上会返回 EIO(bpo-22326); + 本测试只固化「不依赖它」这一约束,真实兼容性需在原 NAS 环境验证。 + """ + target = tmp_path / "fuse_dir" + target.mkdir() + + def mock_temporary_file(*args, **kwargs): + raise OSError(errno.EIO, "Input/output error") + + monkeypatch.setattr("swanlab.sdk.internal.pkg.fs.dir.tempfile.TemporaryFile", mock_temporary_file) + + result = dir.safe_mkdir(target, timeout=1.0) + + assert result == target + # 探针不应留下任何垃圾文件 + assert not list(target.glob(dir.PROBE_PREFIX + "*")) + + +def test_probe_writable_no_leftover(tmp_path: Path): + """探针正常路径:写入成功后应清理本次探针文件""" + dir._probe_writable(tmp_path) + + assert not list(tmp_path.glob(dir.PROBE_PREFIX + "*")) + + +def test_probe_unlink_failure_does_not_fail_writability_probe(monkeypatch, tmp_path: Path): + """unlink 失败只保留本次探针文件,不影响目录可写性判定""" + target = tmp_path / "unlink_fail_dir" + target.mkdir() + + real_unlink = dir.os.unlink + unlink_calls = [] + trace_messages = [] + + def mock_unlink(path, *args, **kwargs): + if Path(path).name.startswith(dir.PROBE_PREFIX): + unlink_calls.append(path) + raise OSError(errno.EIO, "Input/output error") + return real_unlink(path, *args, **kwargs) + + def mock_trace(message, *args, **kwargs): + trace_messages.append((message, kwargs)) + + monkeypatch.setattr("swanlab.sdk.internal.pkg.fs.dir.os.unlink", mock_unlink) + monkeypatch.setattr("swanlab.sdk.internal.pkg.safe.console.trace", mock_trace) + + assert dir.safe_mkdir(target, timeout=5.0) == target + assert len(unlink_calls) == 1 + assert len(list(target.glob(dir.PROBE_PREFIX + "*"))) == 1 + assert len(trace_messages) == 1 + assert "Failed to clean up writability probe file" in trace_messages[0][0] + assert trace_messages[0][1]["write_to_tty"] is False + + +def test_probe_cleanup_failure_does_not_mask_original_error(monkeypatch, tmp_path: Path): + """写入失败且清理(close/unlink)也失败时,原始写入异常不得被覆盖""" + real_unlink = dir.os.unlink + unlink_calls = [] + + def mock_write(fd, data): + raise OSError(errno.EIO, "simulated write failure") + + def mock_unlink(path, *args, **kwargs): + if Path(path).name.startswith(dir.PROBE_PREFIX): + unlink_calls.append(str(path)) + raise OSError(errno.EPERM, "simulated unlink failure") + return real_unlink(path, *args, **kwargs) + + monkeypatch.setattr("swanlab.sdk.internal.pkg.fs.dir.os.write", mock_write) + monkeypatch.setattr("swanlab.sdk.internal.pkg.fs.dir.os.unlink", mock_unlink) + + with pytest.raises(OSError, match="simulated write failure") as exc_info: + dir._probe_writable(tmp_path) + + # 抛出的是原始写入错误而非清理错误,且清理确实被尝试 + assert exc_info.value.errno == errno.EIO + assert len(unlink_calls) == 1 + + def test_timeout_env_invalid_string(monkeypatch): """测试环境变量为非法字符串时,是否能安全回退到默认值 5.0""" monkeypatch.setenv("SWANLAB_FS_TIMEOUT", "invalid_abc123") From 93de72d39f9f64d18d76946cc7010080d627af7a Mon Sep 17 00:00:00 2001 From: CaddiesNew <50736785+Nexisato@users.noreply.github.com> Date: Mon, 21 Sep 2026 22:40:02 +0800 Subject: [PATCH 11/11] fix: filter exp by real time (#1780) Co-authored-by: Kang Li <79990647+SAKURA-CAT@users.noreply.github.com> --- .../internal/core_python/api/experiment.py | 21 ++++--- .../core_python/api/test_experiment.py | 57 +++++++++++++++++-- 2 files changed, 65 insertions(+), 13 deletions(-) diff --git a/swanlab/sdk/internal/core_python/api/experiment.py b/swanlab/sdk/internal/core_python/api/experiment.py index 674104815..49c1af32c 100644 --- a/swanlab/sdk/internal/core_python/api/experiment.py +++ b/swanlab/sdk/internal/core_python/api/experiment.py @@ -47,7 +47,8 @@ def create_or_resume_experiment( :param job_type: 任务类型 :param group: 实验组 :param tags: 实验标签 - :param created_at: 实验创建时间,格式为 ISO 8601 + :param created_at: 实验创建时间,协议保留字段,当前不上报:后端按数据入库时间过滤查询, + 上报的 createdAt 须为指标开始上报的时刻,因此取调用时的当前时间 """ if resume == "must": if run_id is None: @@ -58,11 +59,13 @@ def create_or_resume_experiment( if e.response.status_code == 404 and e.response.reason == "Not Found": raise RuntimeError(f"Experiment {run_id} does not exist in project {project}") labels = [{"name": tag} for tag in tags] if tags else [] - created_at_for_request = created_at.ToDatetime().isoformat() + "Z" + current_ts = Timestamp() + current_ts.GetCurrentTime() + reported_at = current_ts.ToDatetime().isoformat() + "Z" body = { "name": name, "description": description, - "createdAt": created_at_for_request, + "createdAt": reported_at, "colors": [color, color], "labels": labels if len(labels) else None, "job": job_type, @@ -72,11 +75,11 @@ def create_or_resume_experiment( resp = client.post(f"/project/{username}/{project}/experiment", helper.strip_none(body, strip_empty_str=True)) # 200代表实验已存在,开启更新模式 # 201代表实验不存在,新建实验 - # NOTE: 后端返回值没有携带createdAt字段,如果实验不存在则使用前端传入的createdAt字段作为实验创建时间,否则再请求一次获取实验创建时间 + # NOTE: 后端返回值没有携带createdAt字段,如果实验不存在则使用上报的 createdAt(当前时刻)作为实验创建时间,否则再请求一次获取实验创建时间 experiment: InitExperimentType = resp.data is_new_experiment = resp.raw.status_code == 201 if is_new_experiment: - experiment["createdAt"] = created_at_for_request + experiment["createdAt"] = reported_at elif not experiment.get("createdAt"): # 旧后端 POST 响应未携带 createdAt,回退到 GET 获取 exp_resp = client.get(f"/project/{username}/{project}/runs/{experiment['cuid']}") @@ -113,18 +116,22 @@ def stop_experiment(username: str, project: str, experiment_id: str, *, state: R :param project: 所属项目名称 :param experiment_id: 所属实验名称 :param state: 实验状态 - :param finished_at: 实验结束时间 + :param finished_at: 实验结束时间,协议保留字段,当前不上报:后端按数据入库时间过滤查询, + 上报的 finishedAt 须为指标结束上报的时刻,因此取调用时的当前时间 """ this_state: Literal["FINISHED", "CRASHED", "ABORTED"] = "FINISHED" if state == RUN_STATE_CRASHED: this_state = "CRASHED" elif state == RUN_STATE_ABORTED: this_state = "ABORTED" + current_ts = Timestamp() + current_ts.GetCurrentTime() + reported_at = current_ts.ToDatetime().isoformat() + "Z" client.put( f"/project/{username}/{project}/runs/{experiment_id}/state", { "state": this_state, - "finishedAt": finished_at.ToDatetime().isoformat() + "Z", + "finishedAt": reported_at, "from": "sdk", }, ) diff --git a/tests/unit/sdk/internal/core_python/api/test_experiment.py b/tests/unit/sdk/internal/core_python/api/test_experiment.py index 72ab05505..39039316f 100644 --- a/tests/unit/sdk/internal/core_python/api/test_experiment.py +++ b/tests/unit/sdk/internal/core_python/api/test_experiment.py @@ -1,10 +1,11 @@ -from datetime import datetime +from datetime import datetime, timezone from types import SimpleNamespace from unittest.mock import MagicMock import pytest from google.protobuf.timestamp_pb2 import Timestamp +from swanlab.proto.swanlab.run.v1.run_pb2 import RUN_STATE_ABORTED, RUN_STATE_FINISHED from swanlab.sdk.internal.core_python.api import experiment as experiment_api @@ -39,20 +40,25 @@ def _call_create_or_resume(): ) -def test_create_new_experiment_uses_local_created_at(monkeypatch): +def test_create_new_experiment_reports_current_created_at(monkeypatch): + # createdAt 上报的是指标开始上报的时刻(当前时间),而非本地记录的创建时间 post = MagicMock(return_value=_post_resp(_experiment_data(), 201)) get = MagicMock() monkeypatch.setattr(experiment_api.client, "post", post) monkeypatch.setattr(experiment_api.client, "get", get) + before = datetime.now(timezone.utc) experiment, is_new = _call_create_or_resume() assert is_new is True - assert experiment["createdAt"] == "2024-08-01T00:00:00Z" - # 新建实验直接复用请求中的 createdAt,不做额外请求 - get.assert_not_called() assert post.call_args.args[0] == "/project/alice/demo/experiment" - assert post.call_args.args[1]["createdAt"] == "2024-08-01T00:00:00Z" + reported_created_at = post.call_args.args[1]["createdAt"] + assert reported_created_at != "2024-08-01T00:00:00Z" + # 新建实验直接复用上报的 createdAt,不做额外请求 + get.assert_not_called() + assert experiment["createdAt"] == reported_created_at + parsed = datetime.fromisoformat(reported_created_at.replace("Z", "+00:00")) + assert before <= parsed <= datetime.now(timezone.utc) def test_resume_skips_get_when_post_response_contains_created_at(monkeypatch): @@ -93,3 +99,42 @@ def test_resume_raises_when_get_also_missing_created_at(monkeypatch): _call_create_or_resume() get.assert_called_once_with("/project/alice/demo/runs/experiment-id") + + +def _stale_timestamp() -> Timestamp: + ts = Timestamp() + ts.FromDatetime(datetime(2024, 8, 1, tzinfo=timezone.utc)) + return ts + + +def _stop_experiment(monkeypatch) -> MagicMock: + put = MagicMock() + monkeypatch.setattr(experiment_api.client, "put", put) + return put + + +def test_stop_experiment_reports_current_time_as_finished_at(monkeypatch): + # finishedAt 上报的是指标结束上报的时刻(当前时间),而非本地记录的结束时间 + put = _stop_experiment(monkeypatch) + before = datetime.now(timezone.utc) + + experiment_api.stop_experiment( + "alice", "demo", "experiment-id", state=RUN_STATE_FINISHED, finished_at=_stale_timestamp() + ) + + body = put.call_args.args[1] + assert body["state"] == "FINISHED" + assert body["from"] == "sdk" + finished_at = datetime.fromisoformat(body["finishedAt"].replace("Z", "+00:00")) + assert finished_at != datetime(2024, 8, 1, tzinfo=timezone.utc) + assert before <= finished_at <= datetime.now(timezone.utc) + + +def test_stop_experiment_maps_crashed_and_aborted_states(monkeypatch): + put = _stop_experiment(monkeypatch) + + experiment_api.stop_experiment( + "alice", "demo", "experiment-id", state=RUN_STATE_ABORTED, finished_at=_stale_timestamp() + ) + + assert put.call_args.args[1]["state"] == "ABORTED"