From 83aefdabec7ec237b2b7ef324fcfaab234e4bed2 Mon Sep 17 00:00:00 2001 From: Guillaume Date: Wed, 26 Aug 2026 12:20:40 +0200 Subject: [PATCH 01/29] Feat/vendored opencode binary (#299) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * feat(agent): vendored OpenCode binary + auto-install/init and container-tunnel support --------------- - opencode_binary.py: Node-free on-demand provisioning of the OpenCode standalone binary (npm-registry tarball via stdlib), managed per-user cache - resolver prefers managed binary; background auto-install on import/start - `weightslab agent init` CLI; agent-config gating with an info hint (no implicit sign-in) - configurable OpenCode bind host + UI trusted-hosts allowlist so the agent works through a container's published port / SSH tunnel - unit tests + CI agent-smoke job Co-Authored-By: Claude Opus 4.8 (1M context) * docs(agent): Getting Started via `weightslab agent init` + OpenCode env var reference - agent_quickstart: lead with `weightslab agent init` (provisions the Node-free OpenCode binary, then signs in); document --provision-only and the "agent not initialized" info behavior so a new user knows exactly what to do. - configuration: document the OpenCode provisioning/bind env vars — WEIGHTSLAB_OPENCODE_HOST, WEIGHTSLAB_UI_TRUSTED_HOSTS, WEIGHTSLAB_OPENCODE_AUTOINSTALL/AUTODOWNLOAD/VERSION/HOME — incl. the container-behind-a-tunnel setup. --- .github/workflows/ci.yml | 60 +++ docs/agent_quickstart.rst | 32 +- docs/configuration.rst | 63 +++ scripts/ci/agent_smoke.py | 297 ++++++++++++++ tests/test_agent_cli.py | 77 ++++ tests/test_agent_networking.py | 70 ++++ tests/test_opencode_binary.py | 259 ++++++++++++ weightslab/__init__.py | 24 ++ weightslab/cli.py | 148 ++++++- .../PyTorch/wl-classification/main.py | 6 +- weightslab/opencode_binary.py | 372 ++++++++++++++++++ weightslab/opencode_process.py | 55 ++- weightslab/ui/server.py | 93 +++-- 13 files changed, 1510 insertions(+), 46 deletions(-) create mode 100644 scripts/ci/agent_smoke.py create mode 100644 tests/test_agent_cli.py create mode 100644 tests/test_agent_networking.py create mode 100644 tests/test_opencode_binary.py create mode 100644 weightslab/opencode_binary.py diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 454657cb..8526450b 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -224,6 +224,66 @@ jobs: # A per-test timeout guards against any regression that hangs a test. python -m pytest ./tests -v --timeout=600 + # ── Agent smoke test on a pip-installed package ─────────────────────────── + # Proves the Option-2 promise end-to-end: install weightslab into a CLEAN + # virtualenv (from the built wheel, not editable) and confirm the OpenCode + # agent works with NO manual `npm i -g opencode` / `npx` step -- + # 1. the managed OpenCode binary provisions and runs (`--version`), + # 2. `weightslab start` (UI) brings the agent server up even with no + # credential configured (the user can configure it afterwards), and + # 3. `weightslab start example` boots without ever hitting the + # "no opencode/npx" path. + agent-smoke: + needs: [ gate, install ] + if: ${{ needs.gate.outputs.run_ci == 'true' }} + runs-on: ubuntu-latest + timeout-minutes: 30 + name: agent smoke (pip install) + steps: + - name: Checkout repository + uses: actions/checkout@v4 + + - name: Set up Python 3.11 + uses: actions/setup-python@v6 + with: + python-version: '3.11' + + - name: Create clean virtual environment and install from wheel + run: | + python -m venv .venv-smoke + . .venv-smoke/bin/activate + python -m pip install --upgrade pip build + # Build a wheel and install THAT (a real "pip install the package", + # not an editable checkout) so package-data / entry points are exercised + # exactly as an end user would get them. + python -m build --wheel + python -m pip install dist/*.whl --extra-index-url https://download.pytorch.org/whl/cpu + + - name: Agent smoke — provision opencode + weightslab start + env: + WEIGHTSLAB_LOG_LEVEL: INFO + # Force the managed provisioning path (do not depend on the runner's + # preinstalled Node): the standalone binary must run on its own. + WEIGHTSLAB_OPENCODE_AUTODOWNLOAD: '1' + run: | + . .venv-smoke/bin/activate + python scripts/ci/agent_smoke.py start + + - name: Agent smoke — weightslab agent init (CLI, headless) + env: + WEIGHTSLAB_LOG_LEVEL: INFO + WEIGHTSLAB_OPENCODE_AUTODOWNLOAD: '1' + run: | + . .venv-smoke/bin/activate + python scripts/ci/agent_smoke.py cli-init + + - name: Agent smoke — weightslab start example (agent optional, no error) + env: + WEIGHTSLAB_LOG_LEVEL: INFO + run: | + . .venv-smoke/bin/activate + python scripts/ci/agent_smoke.py example + build-and-publish-dev: # Only publish to TestPyPI when pushing to main (not on PRs or dev branch pushes). if: ${{ github.event_name == 'push' && github.ref == 'refs/heads/main' }} diff --git a/docs/agent_quickstart.rst b/docs/agent_quickstart.rst index 50e6da5d..1a243e2c 100644 --- a/docs/agent_quickstart.rst +++ b/docs/agent_quickstart.rst @@ -19,25 +19,39 @@ server. This page is the fastest path from "just installed WeightsLab" to What you need -------------- -- WeightsLab installed (``pip install weightslab``) — this brings the - ``opencode-ai`` bundled binary with it, so there is nothing extra to - install for the agent itself. +- WeightsLab installed (``pip install weightslab``). That's the only install + step: WeightsLab provisions the OpenCode binary itself, on first use, into a + per-user cache — **no Node.js and no manual ``npm``/``opencode`` install + required**. - One set of credentials for a model provider: an OpenRouter API key, an Anthropic key, or a local Ollama install. Pick whichever you already have. -Step 1 — authenticate OpenCode once +Step 1 — initialize the agent once ------------------------------------ The agent's provider and credentials live entirely inside OpenCode, never in -WeightsLab itself. Do this once per machine: +WeightsLab itself. The one-liner below provisions the OpenCode binary (if it +isn't already) and then signs you in — do this once per machine: .. code-block:: bash - opencode auth login + weightslab agent init -Follow the prompts to sign in to OpenRouter, Anthropic, or point it at a -local Ollama endpoint. You can also do this later from the browser, using the -login modal on the Weights Studio landing page — no terminal required. +Follow the prompts to sign in to OpenRouter, Anthropic, or point it at a local +Ollama endpoint. Equivalent alternatives: + +- ``opencode auth login`` — if you prefer to drive OpenCode directly (WeightsLab + installs the binary either way). +- The login modal on the Weights Studio landing page — no terminal required. +- ``weightslab agent init --provision-only`` — headless/CI: just install the + binary, skip the interactive sign-in. + +.. note:: + + You can skip this step and start straight away — if no credential is found, + WeightsLab logs an *info* line ("OpenCode is installed, but the agent is not + initialized yet — run ``weightslab agent init``") and keeps running. The + assistant is optional; nothing else is blocked. Step 2 — start an experiment ------------------------------ diff --git a/docs/configuration.rst b/docs/configuration.rst index 32500c80..15f5d3f2 100644 --- a/docs/configuration.rst +++ b/docs/configuration.rst @@ -815,6 +815,22 @@ server and the backend SDK agent share. These control where it lives. one. Takes precedence over everything else, and configures **both** the UI server and the SDK agent — set it once and the two converge on a single process. + * - ``WEIGHTSLAB_OPENCODE_HOST`` + - ``127.0.0.1`` + - Host the spawned agent server **binds** to. Loopback by default (the + server has filesystem access and must not be reachable off-machine on a + normal local run). Set to ``0.0.0.0`` when running in a container reached + over an SSH tunnel / published port, so the published port can reach it — + the URL handed to the browser stays ``127.0.0.1`` either way. See + :ref:`studio-bridging`. + * - ``WEIGHTSLAB_UI_TRUSTED_HOSTS`` + - *(unset)* + - Comma-separated extra source IPs/CIDRs allowed to call the UI server's + local-only control routes (start agent, notebook, loops). Loopback is + always trusted; behind a tunnel + published port the browser's request + arrives from the container gateway, so set e.g. + ``172.16.0.0/12,192.168.0.0/16`` there (the real trust boundary being the + tunnel + host publishing to ``127.0.0.1``). .. note:: @@ -823,6 +839,53 @@ server and the backend SDK agent share. These control where it lives. browser, this port has to be reachable from the browser's side — see :ref:`studio-bridging`. +Agent installation (OpenCode binary) +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +WeightsLab provisions the OpenCode standalone binary itself — no Node.js +required — the first time it is needed (on ``import weightslab``, +``weightslab start``, ``weightslab start example``, or the first agent use). +It is fetched once into a per-user cache and reused. These control that. + +.. list-table:: + :header-rows: 1 + :widths: 35 15 50 + + * - Variable + - Default + - Description + * - ``WEIGHTSLAB_OPENCODE_AUTOINSTALL`` + - ``1`` + - Auto-install OpenCode in the background (logged) on ``import weightslab`` + / ``weightslab start`` when it isn't already present. Set to ``0`` to + disable the on-import/start install (e.g. air-gapped or CI hosts). + * - ``WEIGHTSLAB_OPENCODE_AUTODOWNLOAD`` + - ``1`` + - Master switch for the network fetch. ``0`` forbids all on-demand + downloads — an already-provisioned binary is still used, but nothing new + is fetched (stricter than ``AUTOINSTALL``, which only gates the + import/start pre-warm). + * - ``WEIGHTSLAB_OPENCODE_VERSION`` + - *(pinned)* + - Override the OpenCode version WeightsLab provisions. Each release pins a + known-good version; set this only to track a different one. + * - ``WEIGHTSLAB_OPENCODE_HOME`` + - *(per-user cache)* + - Directory the managed binary is installed under. Defaults to the + platform cache (``~/.cache/weightslab/opencode`` on Linux, + ``%LOCALAPPDATA%\\weightslab\\opencode`` on Windows, + ``~/Library/Caches/weightslab/opencode`` on macOS). + +.. tip:: + + Sign in once with ``weightslab agent init`` (provisions the binary, then runs + ``opencode auth login``); ``weightslab agent init --provision-only`` just + installs the binary without the interactive sign-in, for headless/CI use. + To uninstall, delete ``$WEIGHTSLAB_OPENCODE_HOME`` (default per-user cache + above) and set ``WEIGHTSLAB_OPENCODE_AUTOINSTALL=0`` (and, to also block the + on-demand fetch, ``WEIGHTSLAB_OPENCODE_AUTODOWNLOAD=0``) so it isn't + re-installed. + Agent Provider Setup ~~~~~~~~~~~~~~~~~~~~ diff --git a/scripts/ci/agent_smoke.py b/scripts/ci/agent_smoke.py new file mode 100644 index 00000000..1dc4afcb --- /dev/null +++ b/scripts/ci/agent_smoke.py @@ -0,0 +1,297 @@ +#!/usr/bin/env python3 +"""End-to-end agent smoke test for a *pip-installed* weightslab. + +Run against a clean environment where weightslab was installed from a wheel and +Node.js is deliberately absent. It proves the Option-2 promise: after +``pip install weightslab`` the OpenCode agent works with no manual install. + +Modes (argv[1]): + provision Provision the managed OpenCode binary and run `--version`. + This is "initializing opencode" with no Node/npx on the box. + start `weightslab start` (the UI): boot it headless, then drive + POST /agent-server/start + GET /agent-server/status. Asserts the + agent server comes up WITHOUT any credential configured -- i.e. an + unconfigured user still gets a running agent they can then configure + (opencode auth login / the landing login modal), rather than a hard + failure. ("allow user to configure agent if not already done") + example `weightslab start example`: boot the bundled training example and + assert it starts cleanly and never hits the "no opencode/npx" path. + all provision, then start, then example. + +Exit code is non-zero on the first failure, with a clear reason. +""" + +import json +import os +import signal +import socket +import subprocess +import sys +import threading +import time +import urllib.request +from pathlib import Path + +OPENCODE_MISSING_MARKER = "Could not provision OpenCode" +NOT_CONFIGURED_MARKER = "not initialized" +INSTALL_MARKERS = ("installing now", "OpenCode installed", "OpenCode ready") +STARTUP_BUDGET = 90.0 # UI / agent readiness +EXAMPLE_MIN_UPTIME = 60.0 # example must survive this long past import +EXAMPLE_BUDGET = 300.0 + + +def log(msg: str) -> None: + print(f"[agent-smoke] {msg}", flush=True) + + +def fail(msg: str) -> "NoReturn": # type: ignore[valid-type] + log(f"FAIL: {msg}") + sys.exit(1) + + +def free_port() -> int: + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s: + s.bind(("127.0.0.1", 0)) + return s.getsockname()[1] + + +def http_get(url: str, timeout: float = 3.0): + req = urllib.request.Request(url, headers={"Origin": "http://127.0.0.1"}) + with urllib.request.urlopen(req, timeout=timeout) as resp: + return resp.status, resp.read() + + +def http_post(url: str, timeout: float = 60.0): + req = urllib.request.Request( + url, data=b"{}", method="POST", + headers={"Content-Type": "application/json", "Origin": "http://127.0.0.1"}, + ) + with urllib.request.urlopen(req, timeout=timeout) as resp: + return resp.status, resp.read() + + +def poll_until(fn, budget: float, what: str): + deadline = time.monotonic() + budget + last = None + while time.monotonic() < deadline: + try: + if fn(): + return True + except Exception as exc: # not up yet + last = exc + time.sleep(1.0) + log(f"timed out waiting for {what} ({last})") + return False + + +class Proc: + """A weightslab subprocess with combined-output capture and tree kill.""" + + def __init__(self, args, env=None): + self.args = args + self.lines = [] + self._proc = subprocess.Popen( + args, + stdout=subprocess.PIPE, stderr=subprocess.STDOUT, + text=True, bufsize=1, + env=env or os.environ.copy(), + start_new_session=True, # own process group so we can kill the tree + ) + threading.Thread(target=self._drain, daemon=True).start() + + def _drain(self): + for line in self._proc.stdout: + self.lines.append(line.rstrip("\n")) + print(f" | {line.rstrip()}", flush=True) + + @property + def output(self) -> str: + return "\n".join(self.lines) + + def alive(self) -> bool: + return self._proc.poll() is None + + def returncode(self): + return self._proc.poll() + + def stop(self): + if self._proc.poll() is not None: + return + try: + os.killpg(os.getpgid(self._proc.pid), signal.SIGTERM) + except Exception: + self._proc.terminate() + try: + self._proc.wait(timeout=15) + except Exception: + try: + os.killpg(os.getpgid(self._proc.pid), signal.SIGKILL) + except Exception: + self._proc.kill() + + +def mode_provision() -> None: + log("provisioning managed OpenCode binary (no Node.js expected on PATH)...") + from weightslab import opencode_binary + + if _which("npx") or _which("node"): + log("note: Node is present; the managed path is still exercised explicitly") + + path = opencode_binary.ensure_managed_binary() + if not path: + fail("ensure_managed_binary() returned None -- provisioning failed") + log(f"managed binary: {path}") + + out = subprocess.run([str(path), "--version"], capture_output=True, text=True, timeout=60) + if out.returncode != 0: + fail(f"`opencode --version` failed (rc={out.returncode}): {out.stderr.strip()}") + log(f"opencode --version -> {out.stdout.strip() or out.stderr.strip()}") + + # And confirm the resolver actually selects it. + from weightslab import opencode_process + argv = opencode_process.resolve_opencode_argv() + if not argv or Path(argv[0]) != Path(path): + fail(f"resolver did not select the managed binary: {argv}") + log("resolver selects the managed binary. provision OK") + + +def _which(name: str): + from shutil import which + return which(name) + + +def mode_start() -> None: + port = free_port() + workspace = Path(os.environ.get("RUNNER_TEMP", "/tmp")) / f"wl-smoke-start-{port}" + workspace.mkdir(parents=True, exist_ok=True) + log(f"launching `weightslab start` on port {port} (workspace {workspace})...") + + proc = Proc([ + "weightslab", "start", str(workspace), + "--no-browser", "--host", "127.0.0.1", "--port", str(port), + ]) + try: + base = f"http://127.0.0.1:{port}" + # /agent-server/status always answers 200 JSON once the HTTP server is + # up, independent of whether the bundled SPA assets are present -- a more + # robust readiness probe than "/" (which 404s on an assets-less build). + if not poll_until(lambda: http_get(base + "/agent-server/status")[0] == 200, + STARTUP_BUDGET, "UI server"): + fail(f"UI did not serve on {base}\n---\n{proc.output}") + log("UI is serving") + + # No credential is configured in CI. The agent server must still come up + # -- the user configures the model/login afterwards. That is the whole + # "configure agent if not already done" guarantee. + status, body = http_post(base + "/agent-server/start", timeout=STARTUP_BUDGET) + payload = json.loads(body or b"{}") + if OPENCODE_MISSING_MARKER in proc.output: + fail("agent start hit the no-opencode path despite a pip install") + if not payload.get("ok"): + fail(f"/agent-server/start not ok: {payload}") + if not payload.get("url"): + fail(f"/agent-server/start returned no url: {payload}") + log(f"agent server up (unconfigured) at {payload['url']}") + + # Status endpoint should now report the running agent. + s_status, s_body = http_get(base + "/agent-server/status") + log(f"/agent-server/status -> {s_status} {s_body[:200]!r}") + log("start mode OK") + finally: + proc.stop() + + +def mode_example() -> None: + """`weightslab start example` is pure training: the agent is lazy and + optional. With NO agent configured it must (c) boot cleanly, (c) never hit + the no-opencode path, and just log an info hint -- no init, no error. We + deliberately do NOT provision opencode here (that would be an init the user + never asked for).""" + workspace = Path(os.environ.get("RUNNER_TEMP", "/tmp")) / "wl-smoke-example" + workspace.mkdir(parents=True, exist_ok=True) + log("launching `weightslab start example` with NO agent configured...") + + # Force the unconfigured state so the info-hint path is what we test. + env = {**os.environ, "WEIGHTSLAB_SUPPRESS_BANNER": "1"} + env.pop("OPENCODE_URL", None) + proc = Proc(["weightslab", "start", "example"], env=env) + try: + start = time.monotonic() + while time.monotonic() - start < EXAMPLE_BUDGET: + if OPENCODE_MISSING_MARKER in proc.output: + fail("example surfaced an opencode error despite the agent being optional") + rc = proc.returncode() + if rc is not None: + if rc == 0: + log("example exited 0 during boot window") + break + fail(f"example exited early with rc={rc}\n---\n{proc.output}") + if time.monotonic() - start >= EXAMPLE_MIN_UPTIME: + log(f"example stayed up {int(EXAMPLE_MIN_UPTIME)}s with no opencode error") + break + time.sleep(2.0) + # Soft checks (may land slightly after boot): the "no init, just info" + # sign-in hint, and the background install being logged. + if NOT_CONFIGURED_MARKER in proc.output: + log("info hint present: user told how to `weightslab agent init`") + else: + log("note: agent-config info hint not observed in captured output") + if any(m in proc.output for m in INSTALL_MARKERS): + log("opencode install was logged during the example run") + else: + log("note: opencode install log not observed (may finish after window)") + log("example mode OK") + finally: + proc.stop() + + +def mode_cli_init() -> None: + """(b) The user can initialize the agent from the CLI. Exercise the + headless path: `weightslab agent init --provision-only` must provision a + working opencode with no Node and no interactive prompt.""" + log("running `weightslab agent init --provision-only`...") + out = subprocess.run( + ["weightslab", "agent", "init", "--provision-only"], + capture_output=True, text=True, timeout=300, + ) + combined = (out.stdout or "") + (out.stderr or "") + print(combined, flush=True) + if out.returncode != 0: + fail(f"`weightslab agent init --provision-only` exited {out.returncode}") + if "OpenCode ready" not in combined: + fail("agent init did not report a provisioned OpenCode binary") + + from weightslab import opencode_binary + path = opencode_binary.find_managed_binary() + if not path: + fail("agent init reported success but no managed binary is present") + ver = subprocess.run([str(path), "--version"], capture_output=True, text=True, timeout=60) + if ver.returncode != 0: + fail(f"provisioned opencode failed `--version` (rc={ver.returncode})") + log(f"cli init OK — opencode {ver.stdout.strip() or ver.stderr.strip()} at {path}") + + +def main() -> None: + mode = sys.argv[1] if len(sys.argv) > 1 else "all" + if mode == "provision": + mode_provision() + elif mode == "start": + mode_provision() + mode_start() + elif mode == "example": + # No provisioning: the example must be clean and agent-free on its own. + mode_example() + elif mode == "cli-init": + mode_cli_init() + elif mode == "all": + mode_provision() + mode_start() + mode_cli_init() + mode_example() + else: + fail(f"unknown mode {mode!r}") + log(f"mode {mode!r}: PASS") + + +if __name__ == "__main__": + main() diff --git a/tests/test_agent_cli.py b/tests/test_agent_cli.py new file mode 100644 index 00000000..7d1f8e91 --- /dev/null +++ b/tests/test_agent_cli.py @@ -0,0 +1,77 @@ +"""Tests for the `weightslab agent` CLI surface and the agent-config gating that +keeps an unconfigured run agent-free and error-free (info hint only).""" + +import os +import unittest +from pathlib import Path +from unittest.mock import patch + +from weightslab import cli + + +class AgentConfiguredTests(unittest.TestCase): + def test_not_configured_when_no_url_no_env_no_auth(self): + with patch.dict(os.environ, {}, clear=False): + os.environ.pop("OPENCODE_URL", None) + with patch.object(cli, "_opencode_auth_paths", + return_value=[Path("/nonexistent/auth.json")]), \ + patch.object(cli, "_agent_env_files", + return_value=[Path("/nonexistent/.env")]): + self.assertFalse(cli.agent_is_configured()) + + def test_configured_when_opencode_url_set(self): + with patch.dict(os.environ, {"OPENCODE_URL": "http://127.0.0.1:4096"}, clear=False): + self.assertTrue(cli.agent_is_configured()) + + def test_configured_when_env_file_present(self): + with patch.dict(os.environ, {}, clear=False): + os.environ.pop("OPENCODE_URL", None) + with patch.object(cli, "_agent_env_files") as envs, \ + patch.object(cli, "_opencode_auth_paths", return_value=[]), \ + patch.object(Path, "is_file", return_value=True): + envs.return_value = [Path("/proj/.env")] + self.assertTrue(cli.agent_is_configured()) + + def test_configured_when_auth_file_present(self): + with patch.dict(os.environ, {}, clear=False): + os.environ.pop("OPENCODE_URL", None) + with patch.object(cli, "_agent_env_files", return_value=[Path("/nonexistent/.env")]), \ + patch.object(cli, "_opencode_auth_paths") as paths, \ + patch.object(Path, "is_file", return_value=True): + paths.return_value = [Path("/whatever/auth.json")] + self.assertTrue(cli.agent_is_configured()) + + +class PrewarmGateTests(unittest.TestCase): + def test_install_always_no_hint_when_configured(self): + # Installing the binary is unconditional; only the sign-in hint is gated. + with patch.object(cli, "agent_is_configured", return_value=True), \ + patch.object(cli, "_prewarm_opencode") as prewarm, \ + patch.object(cli, "_log_agent_config_hint") as hint: + cli._prewarm_opencode_or_hint() + prewarm.assert_called_once() + hint.assert_not_called() + + def test_install_and_hint_when_not_configured(self): + with patch.object(cli, "agent_is_configured", return_value=False), \ + patch.object(cli, "_prewarm_opencode") as prewarm, \ + patch.object(cli, "_log_agent_config_hint") as hint: + cli._prewarm_opencode_or_hint() + prewarm.assert_called_once() + hint.assert_called_once() + + +class AgentParserTests(unittest.TestCase): + def test_agent_init_parsed(self): + args = cli._build_parser().parse_args(["agent", "init", "--provision-only"]) + self.assertEqual(args.command, "agent") + self.assertEqual(args.agent_action, "init") + self.assertTrue(args.provision_only) + + def test_agent_init_defaults_no_provision_only(self): + args = cli._build_parser().parse_args(["agent", "init"]) + self.assertFalse(args.provision_only) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_agent_networking.py b/tests/test_agent_networking.py new file mode 100644 index 00000000..b69de103 --- /dev/null +++ b/tests/test_agent_networking.py @@ -0,0 +1,70 @@ +"""Tests for the container/tunnel-enabling knobs: + * opencode_process.opencode_bind_host() -- what OpenCode binds to + * ui.server._client_is_trusted -- who may hit the local-only control routes + +Both default to the safe, loopback-only behavior; the env overrides only widen +things for the container-behind-a-tunnel deployment. +""" + +import ipaddress +import os +import unittest +from unittest.mock import patch + +from weightslab import opencode_process +from weightslab.ui import server as ui_server + + +class BindHostTests(unittest.TestCase): + def test_default_is_loopback(self): + with patch.dict(os.environ, {}, clear=False): + os.environ.pop(opencode_process.HOST_ENV_VAR, None) + self.assertEqual(opencode_process.opencode_bind_host(), "127.0.0.1") + + def test_env_override(self): + with patch.dict(os.environ, {opencode_process.HOST_ENV_VAR: "0.0.0.0"}, clear=False): + self.assertEqual(opencode_process.opencode_bind_host(), "0.0.0.0") + + def test_blank_env_falls_back_to_default(self): + with patch.dict(os.environ, {opencode_process.HOST_ENV_VAR: " "}, clear=False): + self.assertEqual(opencode_process.opencode_bind_host(), "127.0.0.1") + + +class TrustedClientTests(unittest.TestCase): + def test_loopback_always_trusted(self): + with patch.object(ui_server, "_TRUSTED_CLIENT_NETS", []): + self.assertTrue(ui_server._client_is_trusted("127.0.0.1")) + self.assertTrue(ui_server._client_is_trusted("::1")) + + def test_non_loopback_rejected_by_default(self): + with patch.object(ui_server, "_TRUSTED_CLIENT_NETS", []): + self.assertFalse(ui_server._client_is_trusted("172.17.0.1")) + + def test_trusted_net_allows_gateway(self): + nets = [ipaddress.ip_network("172.16.0.0/12")] + with patch.object(ui_server, "_TRUSTED_CLIENT_NETS", nets): + self.assertTrue(ui_server._client_is_trusted("172.17.0.1")) + self.assertFalse(ui_server._client_is_trusted("8.8.8.8")) + + def test_garbage_addr_is_not_trusted(self): + nets = [ipaddress.ip_network("172.16.0.0/12")] + with patch.object(ui_server, "_TRUSTED_CLIENT_NETS", nets): + self.assertFalse(ui_server._client_is_trusted("not-an-ip")) + + +class TrustedNetsParsingTests(unittest.TestCase): + def test_parse_multiple_and_ignore_invalid(self): + with patch.dict(os.environ, + {"WEIGHTSLAB_UI_TRUSTED_HOSTS": "172.16.0.0/12, bad, 10.1.2.3"}, + clear=False): + nets = ui_server._parse_trusted_client_nets() + self.assertEqual(len(nets), 2) + + def test_empty_env_is_empty_list(self): + with patch.dict(os.environ, {}, clear=False): + os.environ.pop("WEIGHTSLAB_UI_TRUSTED_HOSTS", None) + self.assertEqual(ui_server._parse_trusted_client_nets(), []) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_opencode_binary.py b/tests/test_opencode_binary.py new file mode 100644 index 00000000..18f48f3d --- /dev/null +++ b/tests/test_opencode_binary.py @@ -0,0 +1,259 @@ +"""Tests for weightslab/opencode_binary.py -- the on-demand provisioner that +makes ``pip install weightslab`` ship a working OpenCode with no Node.js. + +Everything here is offline: the one network call (download_managed_binary) is +exercised by patching ``urllib.request.urlopen`` to hand back an in-memory npm +tarball, so the extract/chmod/atomic-rename path is covered without touching the +real registry. Platform selection is exercised by patching the tiny set of +host probes (``sys.platform``, ``platform.machine``, AVX2/musl detection). +""" + +import io +import os +import stat +import tarfile +import tempfile +import threading +import unittest +from pathlib import Path +from unittest.mock import patch + +from weightslab import opencode_binary, opencode_process + + +def _fake_npm_tarball(binary_name: str = "opencode", body: bytes = b"#!/bin/sh\necho ok\n") -> bytes: + """Build an in-memory .tgz laid out like an opencode- npm package + (``package/bin/``), the exact shape _extract_binary looks for.""" + buf = io.BytesIO() + with tarfile.open(fileobj=buf, mode="w:gz") as tar: + info = tarfile.TarInfo(name=f"package/bin/{binary_name}") + info.size = len(body) + info.mode = 0o644 + tar.addfile(info, io.BytesIO(body)) + return buf.getvalue() + + +class ManagedPathTests(unittest.TestCase): + def test_home_env_override_and_version_scoping(self): + with tempfile.TemporaryDirectory() as tmp: + env = {opencode_binary.HOME_ENV_VAR: tmp, opencode_binary.VERSION_ENV_VAR: "9.9.9"} + with patch.dict(os.environ, env, clear=False): + path = opencode_binary.managed_binary_path() + self.assertEqual(Path(path).parent.parent, Path(tmp) / "9.9.9") + self.assertEqual(Path(path).parent.name, "bin") + + def test_pinned_version_env_override(self): + with patch.dict(os.environ, {opencode_binary.VERSION_ENV_VAR: "1.2.3"}, clear=False): + self.assertEqual(opencode_binary.pinned_version(), "1.2.3") + with patch.dict(os.environ, {opencode_binary.VERSION_ENV_VAR: ""}, clear=False): + self.assertEqual(opencode_binary.pinned_version(), + opencode_binary.DEFAULT_OPENCODE_VERSION) + + def test_autodownload_toggle(self): + for val, expected in [("0", False), ("false", False), ("no", False), + ("off", False), ("1", True), ("", True)]: + with patch.dict(os.environ, {opencode_binary.AUTODOWNLOAD_ENV_VAR: val}, clear=False): + self.assertEqual(opencode_binary.autodownload_enabled(), expected) + + +class CandidatePackageTests(unittest.TestCase): + def _candidates(self, plat, machine, avx2, musl): + with patch.object(opencode_binary.sys, "platform", plat), \ + patch.object(opencode_binary.platform, "machine", return_value=machine), \ + patch.object(opencode_binary, "_supports_avx2", return_value=avx2), \ + patch.object(opencode_binary, "_is_musl", return_value=musl): + return opencode_binary.candidate_packages() + + def test_linux_x64_avx2_glibc(self): + got = self._candidates("linux", "x86_64", avx2=True, musl=False) + self.assertEqual(got[0], "opencode-linux-x64") + self.assertIn("opencode-linux-x64-baseline", got) + + def test_linux_x64_no_avx2_prefers_baseline(self): + got = self._candidates("linux", "x86_64", avx2=False, musl=False) + self.assertEqual(got[0], "opencode-linux-x64-baseline") + + def test_linux_musl_prefers_musl(self): + got = self._candidates("linux", "x86_64", avx2=True, musl=True) + self.assertEqual(got[0], "opencode-linux-x64-musl") + + def test_linux_arm64(self): + got = self._candidates("linux", "aarch64", avx2=False, musl=False) + self.assertEqual(got[0], "opencode-linux-arm64") + + def test_darwin_arm64(self): + got = self._candidates("darwin", "arm64", avx2=False, musl=False) + self.assertEqual(got, ["opencode-darwin-arm64"]) + + def test_windows_x64_avx2(self): + got = self._candidates("win32", "AMD64", avx2=True, musl=False) + self.assertEqual(got[0], "opencode-windows-x64") + + +class FindAndEnsureTests(unittest.TestCase): + def setUp(self): + self._tmp = tempfile.TemporaryDirectory() + self.addCleanup(self._tmp.cleanup) + self._env = patch.dict( + os.environ, + {opencode_binary.HOME_ENV_VAR: self._tmp.name, + opencode_binary.VERSION_ENV_VAR: "1.2.3"}, + clear=False, + ) + self._env.start() + self.addCleanup(self._env.stop) + + def _install_fake(self): + path = opencode_binary.managed_binary_path() + path.parent.mkdir(parents=True, exist_ok=True) + path.write_text("#!/bin/sh\n") + os.chmod(path, os.stat(path).st_mode | stat.S_IXUSR) + return path + + def test_find_missing_returns_none(self): + self.assertIsNone(opencode_binary.find_managed_binary()) + + def test_find_present_returns_path(self): + path = self._install_fake() + self.assertEqual(opencode_binary.find_managed_binary(), path) + + def test_ensure_returns_existing_without_download(self): + path = self._install_fake() + with patch.object(opencode_binary, "download_managed_binary") as dl: + self.assertEqual(opencode_binary.ensure_managed_binary(), path) + dl.assert_not_called() + + def test_ensure_respects_autodownload_disabled(self): + with patch.dict(os.environ, {opencode_binary.AUTODOWNLOAD_ENV_VAR: "0"}, clear=False): + with patch.object(opencode_binary, "download_managed_binary") as dl: + self.assertIsNone(opencode_binary.ensure_managed_binary()) + dl.assert_not_called() + + def test_ensure_downloads_when_missing(self): + sentinel = self._tmp.name + "/sentinel" + with patch.object(opencode_binary, "download_managed_binary", return_value=sentinel) as dl: + self.assertEqual(opencode_binary.ensure_managed_binary(), sentinel) + dl.assert_called_once() + + +class DownloadTests(unittest.TestCase): + def setUp(self): + self._tmp = tempfile.TemporaryDirectory() + self.addCleanup(self._tmp.cleanup) + self._env = patch.dict( + os.environ, + {opencode_binary.HOME_ENV_VAR: self._tmp.name, + opencode_binary.VERSION_ENV_VAR: "1.2.3"}, + clear=False, + ) + self._env.start() + self.addCleanup(self._env.stop) + + def test_download_extracts_and_marks_executable(self): + tgz = _fake_npm_tarball(binary_name=opencode_binary._binary_filename()) + + def fake_urlopen(req, timeout=None): + return io.BytesIO(tgz) + + with patch.object(opencode_binary.urllib.request, "urlopen", side_effect=fake_urlopen): + path = opencode_binary.download_managed_binary() + + self.assertIsNotNone(path) + self.assertTrue(Path(path).is_file()) + self.assertTrue(os.access(str(path), os.X_OK)) + self.assertEqual(Path(path).read_bytes()[:2], b"#!") + + def test_download_all_candidates_fail_returns_none(self): + def boom(req, timeout=None): + raise OSError("network down") + + with patch.object(opencode_binary.urllib.request, "urlopen", side_effect=boom): + self.assertIsNone(opencode_binary.download_managed_binary()) + + +class BackgroundInstallTests(unittest.TestCase): + def setUp(self): + # Reset the once-per-process guard so each test starts clean. + opencode_binary._bg_started = False + self.addCleanup(setattr, opencode_binary, "_bg_started", False) + + def test_noop_when_already_installed(self): + with patch.object(opencode_binary, "find_managed_binary", return_value=Path("/x/opencode")), \ + patch.object(opencode_binary, "download_managed_binary") as dl: + opencode_binary.ensure_managed_binary_in_background(reason="test") + dl.assert_not_called() + + def test_noop_when_autodownload_disabled(self): + with patch.dict(os.environ, {opencode_binary.AUTODOWNLOAD_ENV_VAR: "0"}, clear=False), \ + patch.object(opencode_binary, "find_managed_binary", return_value=None), \ + patch.object(opencode_binary, "download_managed_binary") as dl: + opencode_binary.ensure_managed_binary_in_background(reason="test") + dl.assert_not_called() + + def test_downloads_in_background_when_missing(self): + done = threading.Event() + + def fake_download(version=None): + done.set() + return Path("/x/opencode") + + with patch.object(opencode_binary, "find_managed_binary", return_value=None), \ + patch.object(opencode_binary, "download_managed_binary", side_effect=fake_download): + opencode_binary.ensure_managed_binary_in_background(reason="test") + self.assertTrue(done.wait(timeout=5), "background download did not run") + + def test_only_first_call_starts(self): + with patch.object(opencode_binary, "find_managed_binary", return_value=None), \ + patch.object(opencode_binary, "download_managed_binary", + return_value=Path("/x/opencode")) as dl: + opencode_binary.ensure_managed_binary_in_background(reason="a") + opencode_binary.ensure_managed_binary_in_background(reason="b") + # second call must be a no-op regardless of thread timing + import time as _t + _t.sleep(0.5) + self.assertLessEqual(dl.call_count, 1) + + +class ResolverPrecedenceTests(unittest.TestCase): + """opencode_process.resolve_opencode_argv order: + managed-present -> PATH -> managed-download -> npx -> None.""" + + def test_managed_present_wins(self): + with patch.object(opencode_process.opencode_binary, "find_managed_binary", + return_value=Path("/mgd/opencode")), \ + patch.object(opencode_process.shutil, "which", return_value="/usr/bin/opencode"): + self.assertEqual(opencode_process.resolve_opencode_argv(), ["/mgd/opencode"]) + + def test_path_used_before_download(self): + with patch.object(opencode_process.opencode_binary, "find_managed_binary", return_value=None), \ + patch.object(opencode_process.opencode_binary, "ensure_managed_binary") as ensure, \ + patch.object(opencode_process.shutil, "which", + side_effect=lambda n: "/usr/bin/opencode" if n == "opencode" else None): + self.assertEqual(opencode_process.resolve_opencode_argv(), ["/usr/bin/opencode"]) + ensure.assert_not_called() + + def test_download_when_no_path(self): + with patch.object(opencode_process.opencode_binary, "find_managed_binary", return_value=None), \ + patch.object(opencode_process.opencode_binary, "ensure_managed_binary", + return_value=Path("/mgd/opencode")), \ + patch.object(opencode_process.shutil, "which", return_value=None): + self.assertEqual(opencode_process.resolve_opencode_argv(), ["/mgd/opencode"]) + + def test_npx_last_resort(self): + def which(name): + return "/usr/bin/npx" if name == "npx" else None + with patch.object(opencode_process.opencode_binary, "find_managed_binary", return_value=None), \ + patch.object(opencode_process.opencode_binary, "ensure_managed_binary", return_value=None), \ + patch.object(opencode_process.shutil, "which", side_effect=which): + self.assertEqual(opencode_process.resolve_opencode_argv(), + ["/usr/bin/npx", "--yes", "opencode-ai@latest"]) + + def test_none_when_nothing_available(self): + with patch.object(opencode_process.opencode_binary, "find_managed_binary", return_value=None), \ + patch.object(opencode_process.opencode_binary, "ensure_managed_binary", return_value=None), \ + patch.object(opencode_process.shutil, "which", return_value=None): + self.assertIsNone(opencode_process.resolve_opencode_argv()) + + +if __name__ == "__main__": + unittest.main() diff --git a/weightslab/__init__.py b/weightslab/__init__.py index 44323142..759de5ea 100644 --- a/weightslab/__init__.py +++ b/weightslab/__init__.py @@ -96,6 +96,30 @@ def __dir__(): if os.getenv('WEIGHTSLAB_SUPPRESS_BANNER', '0') != '1': logger.info(_BANNER) +# Auto-install the OpenCode agent binary on first import, in the background and +# logged, so `import weightslab` / `weightslab start` / `weightslab start example` +# all leave the agent ready with no manual step (installation only — signing in +# stays opt-in via `weightslab agent init`). Guards: main process only (never +# DataLoader workers); skipped under pytest and when +# WEIGHTSLAB_OPENCODE_AUTOINSTALL is falsey or the download is disabled. The +# call itself no-ops instantly when the binary is already present. +def _autoinstall_opencode_on_import(): + import sys as _sys + if os.environ.get('WEIGHTSLAB_OPENCODE_AUTOINSTALL', '1').strip().lower() in {'0', 'false', 'no', 'off'}: + return + if 'pytest' in _sys.modules or os.environ.get('PYTEST_CURRENT_TEST'): + return + try: + from weightslab import opencode_binary + opencode_binary.ensure_managed_binary_in_background( + reason="import weightslab", logger=logger) + except Exception as _exc: # pragma: no cover - best-effort + logger.debug("OpenCode auto-install skipped: %s", _exc) + + +if _IS_MAIN_PROCESS: + _autoinstall_opencode_on_import() + grpc_tls_enabled = os.environ.get('GRPC_TLS_ENABLED', 'true').lower() == 'true' if _IS_MAIN_PROCESS and grpc_tls_enabled and os.environ.get('WEIGHTSLAB_SKIP_SECURE_INIT', 'false').lower() != 'true': try: diff --git a/weightslab/cli.py b/weightslab/cli.py index a70d6b76..d4e5adb4 100644 --- a/weightslab/cli.py +++ b/weightslab/cli.py @@ -602,6 +602,13 @@ def example_start(args): logger.info(f"Starting the WeightsLab {label} ({kind}) example...") logger.info(f" {main_py}") + # Install OpenCode in the background (logged) so it's ready if the user opens + # the chat; the example itself is pure training and never needs it to run. + # Configuration (sign-in) stays optional -- surface it as info, never an + # error, so a run with no agent configured doesn't look like it failed. + _prewarm_opencode() + if not agent_is_configured(): + _log_agent_config_hint() logger.info("In another terminal, launch the UI with: weightslab start") logger.info("Then open the URL printed by `weightslab start` — stop the example with Ctrl+C.") if not _CERTS_DIR_IN_ORIGINAL_ENV: @@ -777,6 +784,83 @@ def _print_experiment_guidance(experiment_dir: Path) -> None: logger.info("=" * 60) +def _opencode_auth_paths() -> "list[Path]": + """Candidate locations of OpenCode's own credential store (written by + `opencode auth login`). Existence of any means the agent has been + initialized on this machine.""" + candidates = [] + xdg = os.environ.get("XDG_DATA_HOME", "").strip() + if xdg: + candidates.append(Path(xdg) / "opencode" / "auth.json") + if _is_windows(): + for var in ("APPDATA", "LOCALAPPDATA"): + base = os.environ.get(var) + if base: + candidates.append(Path(base) / "opencode" / "auth.json") + candidates.append(Path.home() / ".local" / "share" / "opencode" / "auth.json") + return candidates + + +def _agent_env_files() -> "list[Path]": + """Candidate user-provided agent env files (.env). Mirrors the paths + agent.py's _load_config actually loads, so "found here" == "used there".""" + pkg_dir = Path(__file__).resolve().parent # .../weightslab + return [Path.cwd() / ".env", pkg_dir / ".env", pkg_dir.parent / ".env"] + + +def agent_is_configured() -> bool: + """True if the agent has been initialized -- i.e. an env file / credential is + present so it can be used without any further step: + * OPENCODE_URL set (explicit server), or + * a user `.env` present (agent.py loads it), or + * a prior `opencode auth login` (its credential store exists). + + Deliberately does NOT count agent_config.yaml: it ships in the repo/wheel and + only carries URL/model *defaults*, not a credential, so counting it would + make the "not initialized" hint never fire. + """ + if os.environ.get("OPENCODE_URL", "").strip(): + return True + if any(p.is_file() for p in _agent_env_files()): + return True + return any(p.is_file() for p in _opencode_auth_paths()) + + +def _log_agent_config_hint() -> None: + """Tell the user OpenCode is installed but the agent still needs a one-time + sign-in -- an INFO note, not an error. The assistant is optional; a run that + never uses it must not look broken.""" + logger.info( + "OpenCode is installed, but the agent is not initialized yet — run " + "`weightslab agent init` to sign in (or set OPENCODE_URL / add a .env at " + "your project root). The assistant is optional; continuing without it." + ) + + +def _prewarm_opencode() -> None: + """Install OpenCode in the background if not already present (logged). + + Delegates to the shared, once-per-process installer so the ~180 MB first-run + download never blocks the UI, and repeated launches don't re-download. + """ + from weightslab import opencode_binary + opencode_binary.ensure_managed_binary_in_background(reason="weightslab start", logger=logger) + + +def _prewarm_opencode_or_hint() -> None: + """Install OpenCode up front (background, logged) so the agent is ready, and + -- separately -- hint how to sign in when nothing is configured yet. + + Installing the binary is unconditional: it is a free, one-time fetch and is + what makes `weightslab start` leave the agent usable. Signing in stays + opt-in (`weightslab agent init`); we only *hint* at it, never do it + implicitly -- that is the "no agent env found -> no init, just info" rule. + """ + _prewarm_opencode() + if not agent_is_configured(): + _log_agent_config_hint() + + def ui_start_native(args): """`weightslab start`: launch the Weights Studio UI natively (no Docker). @@ -810,6 +894,12 @@ def ui_start_native(args): os.environ["WL_LAST_EXPERIMENT_DIR"] = str(experiment_dir) _print_experiment_guidance(experiment_dir) + # If the agent has been initialized, provision OpenCode up front (in the + # background) so it is ready the moment the user opens the chat. If it has + # NOT been initialized, do nothing but log how to enable it -- an + # unconfigured run stays agent-free and error-free. + _prewarm_opencode_or_hint() + ui_host = getattr(args, "host", None) or os.getenv("WEIGHTSLAB_UI_HOST", "0.0.0.0") preferred_ui_port, ui_port_source = _resolve_ui_port(args) if ui_port_source == "default": @@ -927,6 +1017,46 @@ def _add_example_kind_flags(p: argparse.ArgumentParser) -> None: p.set_defaults(example_kind=_DEFAULT_EXAMPLE) +def agent_init(args): + """`weightslab agent init`: initialize the AI assistant. + + Provisions the OpenCode binary (Node-free, via weightslab.opencode_binary) + and then runs `opencode auth login` so the user signs in once. This is the + explicit opt-in that `weightslab start` / `start example` only *hint* at when + no agent is configured -- nothing here happens implicitly. + """ + from weightslab import opencode_binary + + logger.info("Provisioning the OpenCode agent binary (no Node.js required)...") + path = opencode_binary.ensure_managed_binary() + if not path: + logger.error( + "Could not provision OpenCode (offline?). Retry with network access, " + "or install Node.js 20+ and `npm i -g opencode-ai`." + ) + sys.exit(1) + logger.info(f"OpenCode ready: {path}") + + if getattr(args, "provision_only", False): + logger.info("Provision-only: skipping interactive sign-in.") + logger.info(f"Sign in later with: weightslab agent init (or: {path} auth login)") + return + + logger.info("Launching `opencode auth login` — follow the prompts to sign in.") + try: + rc = subprocess.run([str(path), "auth", "login"]).returncode + except KeyboardInterrupt: + logger.info("Sign-in cancelled. Re-run `weightslab agent init` anytime.") + return + if rc != 0: + logger.warning( + f"`opencode auth login` exited with code {rc}. " + "Re-run `weightslab agent init` to try again." + ) + sys.exit(rc) + logger.info("Agent initialized. The assistant is now available in `weightslab start`.") + + def _build_parser() -> argparse.ArgumentParser: """Build the top-level argument parser (banner + detailed command reference). @@ -945,7 +1075,7 @@ def _build_parser() -> argparse.ArgumentParser: epilog=_EPILOG, formatter_class=argparse.RawDescriptionHelpFormatter, ) - sub = parser.add_subparsers(dest="command", metavar="{se,start,cli,tunnel,export,help}") + sub = parser.add_subparsers(dest="command", metavar="{se,start,cli,tunnel,export,agent,help}") # weightslab se [--force-certs] [certs_dir] se_parser = sub.add_parser("se", help="Set up the secure environment (TLS certs + gRPC auth token)") @@ -1035,6 +1165,17 @@ def _build_parser() -> argparse.ArgumentParser: "start", help="Start a bundled PyTorch example (default: classification)") _add_example_kind_flags(example_alias_start) + # weightslab agent init [--provision-only] + agent_parser = sub.add_parser( + "agent", help="Manage the AI assistant (OpenCode), e.g. `weightslab agent init`") + agent_sub = agent_parser.add_subparsers(dest="agent_action", metavar="{init}") + agent_init_parser = agent_sub.add_parser( + "init", help="Provision OpenCode and sign in so the assistant is ready to use") + agent_init_parser.add_argument( + "--provision-only", action="store_true", + help="Only download/verify the OpenCode binary; skip the interactive " + "sign-in (for headless/CI environments).") + sub.add_parser("help", help="Show this help message") return parser @@ -1070,6 +1211,11 @@ def main(): elif args.command == "example": # Alias for `start example` — tolerate the swapped subcommand order. example_start(args) + elif args.command == "agent": + if getattr(args, "agent_action", None) == "init": + agent_init(args) + else: + _build_parser().parse_args(["agent", "--help"]) else: parser.print_help() diff --git a/weightslab/examples/PyTorch/wl-classification/main.py b/weightslab/examples/PyTorch/wl-classification/main.py index 7f095c25..dc64f510 100644 --- a/weightslab/examples/PyTorch/wl-classification/main.py +++ b/weightslab/examples/PyTorch/wl-classification/main.py @@ -385,9 +385,9 @@ def test(loader, model, criterion_mlt, metric_mlt, device, test_loader_len): else: train_range = itertools.count() - # # ============= - # # Training Loop - # wl.start_training(timeout=3) # Blocks and keeps the main thread alive while background services run. Optionally set a timeout (seconds) to auto-stop. + # ============= + # Training Loop + wl.start_training(timeout=3) # Blocks and keeps the main thread alive while background services run. Optionally set a timeout (seconds) to auto-stop. train_loss = None test_loss, test_metric = None, None diff --git a/weightslab/opencode_binary.py b/weightslab/opencode_binary.py new file mode 100644 index 00000000..35c2cc9c --- /dev/null +++ b/weightslab/opencode_binary.py @@ -0,0 +1,372 @@ +"""Self-contained provisioning of a standalone OpenCode binary. + +Why this exists +--------------- +``weightslab`` drives a local ``opencode serve`` process (see +``opencode_process.py``). Historically the only ways to get that binary were a +global ``npm i -g opencode-ai`` or the ``npx --yes`` fallback -- both of which +need Node.js on the machine. That makes a plain ``pip install weightslab`` in a +clean environment *not* enough: the agent silently fails with "Could not find +`opencode` or `npx`" until the user installs Node and OpenCode by hand. + +This module removes that manual step without bloating the wheel. OpenCode ships +its ~180 MB standalone binaries inside platform-specific npm packages +(``opencode-linux-x64``, ``opencode-darwin-arm64``, ...), each downloadable as a +plain gzip tarball from the npm registry -- no Node needed to fetch or unpack, +only ``urllib`` + ``tarfile`` from the stdlib. So instead of vendoring a +180 MB-per-platform binary into every wheel (which would make wheels huge and +platform-locked), we fetch the *correct* binary on demand into a per-user cache +and reuse it forever after. ``resolve_opencode_argv`` prefers this managed +binary, so after a clean ``pip install`` the agent "just works". + +The platform/arch/musl/AVX2 selection logic mirrors ``opencode-ai``'s own +``postinstall.mjs`` so we pick exactly the package its installer would have. +""" + +from __future__ import annotations + +import logging +import os +import platform +import shutil +import stat +import subprocess +import sys +import tarfile +import tempfile +import threading +import urllib.request +from pathlib import Path +from typing import List, Optional + +_LOGGER = logging.getLogger(__name__) + +# The OpenCode version this weightslab release pins. Kept explicit (not +# "@latest") so a given weightslab build always provisions a known-good, +# tested OpenCode -- reproducible installs, no surprise upgrade mid-release. +# Override with WEIGHTSLAB_OPENCODE_VERSION to track a different one. +DEFAULT_OPENCODE_VERSION = "1.18.23" + +VERSION_ENV_VAR = "WEIGHTSLAB_OPENCODE_VERSION" +HOME_ENV_VAR = "WEIGHTSLAB_OPENCODE_HOME" +# Set to "0"/"false"/"no" to forbid the on-demand download (air-gapped hosts, +# CI that must stay offline). find_managed_binary() still returns an already +# provisioned binary; only the network fetch is suppressed. +AUTODOWNLOAD_ENV_VAR = "WEIGHTSLAB_OPENCODE_AUTODOWNLOAD" + +_REGISTRY = "https://registry.npmjs.org" +# Generous: a cold fetch pulls a ~180 MB tarball over the public registry. +_DOWNLOAD_TIMEOUT = 180.0 + + +def pinned_version() -> str: + """The OpenCode version to provision (env override wins).""" + return os.environ.get(VERSION_ENV_VAR, "").strip() or DEFAULT_OPENCODE_VERSION + + +def autodownload_enabled() -> bool: + raw = os.environ.get(AUTODOWNLOAD_ENV_VAR, "").strip().lower() + if raw in {"0", "false", "no", "off"}: + return False + return True + + +def _cache_root() -> Path: + """Per-user cache directory the managed binary lives under. + + Honours WEIGHTSLAB_OPENCODE_HOME, then the platform-conventional cache + location, so provisioning survives across virtualenvs (the binary is a + property of the machine, not of one env) and never needs write access to + the -- possibly read-only -- site-packages tree. + """ + override = os.environ.get(HOME_ENV_VAR, "").strip() + if override: + return Path(override).expanduser() + + if sys.platform == "win32": + base = os.environ.get("LOCALAPPDATA") or os.environ.get("APPDATA") + root = Path(base) if base else Path.home() / "AppData" / "Local" + return root / "weightslab" / "opencode" + if sys.platform == "darwin": + return Path.home() / "Library" / "Caches" / "weightslab" / "opencode" + xdg = os.environ.get("XDG_CACHE_HOME", "").strip() + root = Path(xdg) if xdg else Path.home() / ".cache" + return root / "weightslab" / "opencode" + + +def _binary_filename() -> str: + # OpenCode names the extracted binary opencode.exe on Windows, opencode + # elsewhere (postinstall.mjs's sourceBinary). + return "opencode.exe" if sys.platform == "win32" else "opencode" + + +def managed_binary_path(version: Optional[str] = None) -> Path: + """Where the managed binary for ``version`` is (or would be) installed. + + Version-scoped so bumping DEFAULT_OPENCODE_VERSION provisions cleanly + alongside the old one instead of clobbering a binary another env still uses. + """ + version = version or pinned_version() + return _cache_root() / version / "bin" / _binary_filename() + + +def _norm_platform() -> str: + return {"darwin": "darwin", "linux": "linux", "win32": "windows"}.get( + sys.platform, sys.platform + ) + + +def _norm_arch() -> str: + machine = platform.machine().lower() + if machine in {"x86_64", "amd64", "x64"}: + return "x64" + if machine in {"arm64", "aarch64"}: + return "arm64" + if machine.startswith("arm"): + return "arm" + return machine + + +def _supports_avx2() -> bool: + """AVX2 probe, x64 only -- mirrors postinstall.mjs. Non-AVX2 x64 CPUs need + the ``-baseline`` build; getting this wrong yields an illegal-instruction + crash at first run, so we default to the safe (baseline-preferred) answer + whenever detection is uncertain.""" + if _norm_arch() != "x64": + return False + system = _norm_platform() + try: + if system == "linux": + with open("/proc/cpuinfo", "r", encoding="utf-8", errors="ignore") as fh: + return " avx2 " in (" " + fh.read().lower() + " ") + if system == "darwin": + out = subprocess.run( + ["sysctl", "-n", "hw.optional.avx2_0"], + capture_output=True, text=True, timeout=1.5, + ) + return out.returncode == 0 and out.stdout.strip() == "1" + if system == "windows": + # IsProcessorFeaturePresent(40) == PF_AVX2_INSTRUCTIONS_AVAILABLE. + ps = ( + '(Add-Type -MemberDefinition "[DllImport(\\"kernel32.dll\\")] ' + 'public static extern bool IsProcessorFeaturePresent(int f);" ' + "-Name K -Namespace W -PassThru)::IsProcessorFeaturePresent(40)" + ) + for exe in ("powershell.exe", "pwsh.exe", "pwsh", "powershell"): + if not shutil.which(exe): + continue + out = subprocess.run( + [exe, "-NoProfile", "-NonInteractive", "-Command", ps], + capture_output=True, text=True, timeout=3.0, + ) + if out.returncode == 0: + return out.stdout.strip().lower() in {"true", "1"} + except Exception: # pragma: no cover - detection is best-effort + return False + return False + + +def _is_musl() -> bool: + if _norm_platform() != "linux": + return False + try: + if Path("/etc/alpine-release").exists(): + return True + except Exception: # pragma: no cover - filesystem probe blocked + pass + try: + out = subprocess.run(["ldd", "--version"], capture_output=True, text=True) + return "musl" in (out.stdout + out.stderr).lower() + except Exception: # pragma: no cover - ldd absent + return False + + +def candidate_packages() -> List[str]: + """Ordered npm package names to try for this host, most-preferred first. + + Mirrors opencode-ai/postinstall.mjs's ``packageNames()`` so we resolve the + same artifact its own installer would, including the -musl and -baseline + fallbacks. The list is ordered, not singular, precisely so a wrong AVX2/musl + guess degrades to a working build rather than a hard failure. + """ + system = _norm_platform() + arch = _norm_arch() + base = f"opencode-{system}-{arch}" + baseline = arch == "x64" and not _supports_avx2() + + if system == "linux": + if _is_musl(): + if arch == "x64": + return ( + [f"{base}-baseline-musl", f"{base}-musl", f"{base}-baseline", base] + if baseline + else [f"{base}-musl", f"{base}-baseline-musl", base, f"{base}-baseline"] + ) + return [f"{base}-musl", base] + if arch == "x64": + return ( + [f"{base}-baseline", base, f"{base}-baseline-musl", f"{base}-musl"] + if baseline + else [base, f"{base}-baseline", f"{base}-musl", f"{base}-baseline-musl"] + ) + return [base, f"{base}-musl"] + + if arch == "x64": + return [f"{base}-baseline", base] if baseline else [base, f"{base}-baseline"] + return [base] + + +def _tarball_url(pkg: str, version: str) -> str: + # Standard unscoped-package layout on the npm registry. + return f"{_REGISTRY}/{pkg}/-/{pkg}-{version}.tgz" + + +def _extract_binary(tgz_path: Path, dest: Path) -> bool: + """Extract ``package/bin/`` from an npm tarball to ``dest``. + + Writes to a sibling temp file and atomically renames, so a concurrent + reader never sees a half-written binary and two racing provisioners can't + corrupt each other's output. + """ + wanted = f"bin/{_binary_filename()}" + dest.parent.mkdir(parents=True, exist_ok=True) + with tarfile.open(tgz_path, "r:gz") as tar: + member = next( + (m for m in tar.getmembers() if m.isfile() and m.name.replace("\\", "/").endswith(wanted)), + None, + ) + if member is None: + _LOGGER.warning("OpenCode tarball %s has no %s", tgz_path.name, wanted) + return False + src = tar.extractfile(member) + if src is None: # pragma: no cover - defensive + return False + fd, tmp_name = tempfile.mkstemp(dir=str(dest.parent), prefix=".opencode-", suffix=".part") + tmp = Path(tmp_name) + try: + with os.fdopen(fd, "wb") as out: + shutil.copyfileobj(src, out, length=1024 * 1024) + mode = os.stat(tmp).st_mode + os.chmod(tmp, mode | stat.S_IXUSR | stat.S_IXGRP | stat.S_IXOTH) + os.replace(tmp, dest) + finally: + if tmp.exists(): + tmp.unlink() + return True + + +def download_managed_binary(version: Optional[str] = None) -> Optional[Path]: + """Fetch and install the managed OpenCode binary. Returns its path or None. + + Tries each candidate package in turn; the first that both downloads and + yields the wanted binary wins. Never raises: a failed provision degrades to + ``None`` so the caller can fall back to PATH/npx rather than crash. + """ + version = version or pinned_version() + dest = managed_binary_path(version) + packages = candidate_packages() + _LOGGER.info( + "OpenCode: provisioning managed binary %s (%s) -> %s", + version, packages[0] if packages else "?", dest, + ) + for pkg in packages: + url = _tarball_url(pkg, version) + try: + with tempfile.NamedTemporaryFile(delete=False, suffix=".tgz") as tmp: + tgz = Path(tmp.name) + req = urllib.request.Request(url, headers={"User-Agent": "weightslab-opencode-provisioner"}) + with urllib.request.urlopen(req, timeout=_DOWNLOAD_TIMEOUT) as resp: + if getattr(resp, "status", 200) != 200: + continue + with open(tgz, "wb") as fh: + shutil.copyfileobj(resp, fh, length=1024 * 1024) + ok = _extract_binary(tgz, dest) + if ok: + _LOGGER.info("OpenCode: installed %s from %s", dest, pkg) + return dest + except Exception as exc: # try the next candidate package + _LOGGER.debug("OpenCode: candidate %s failed (%s)", pkg, exc) + continue + finally: + try: + tgz.unlink() + except Exception: + pass + _LOGGER.warning( + "OpenCode: could not provision a managed binary for %s/%s (version %s).", + _norm_platform(), _norm_arch(), version, + ) + return None + + +def _looks_runnable(path: Path) -> bool: + try: + return path.is_file() and os.access(str(path), os.X_OK) + except Exception: # pragma: no cover - defensive + return False + + +def find_managed_binary(version: Optional[str] = None) -> Optional[Path]: + """Return an already-provisioned managed binary, or None. No network.""" + path = managed_binary_path(version) + return path if _looks_runnable(path) else None + + +def ensure_managed_binary(version: Optional[str] = None, + auto_download: Optional[bool] = None) -> Optional[Path]: + """Return a usable managed binary, downloading it once if needed. + + ``auto_download`` defaults to the WEIGHTSLAB_OPENCODE_AUTODOWNLOAD env + setting. Returns None (never raises) when no managed binary can be made + available, so callers can fall back to PATH/npx. + """ + existing = find_managed_binary(version) + if existing: + return existing + if auto_download is None: + auto_download = autodownload_enabled() + if not auto_download: + return None + return download_managed_binary(version) + + +# Guard so many callers (import hook + `weightslab start` + `start example`) at +# most spawn ONE background install per process, rather than racing downloads. +_bg_lock = threading.Lock() +_bg_started = False + + +def ensure_managed_binary_in_background(reason: str = "", logger: Optional[logging.Logger] = None) -> None: + """Install OpenCode in a daemon thread if it isn't already present, logging + the install. Idempotent, best-effort, and non-blocking -- the caller (a CLI + launch or ``import weightslab``) never waits on the ~180 MB download. + + Respects WEIGHTSLAB_OPENCODE_AUTODOWNLOAD. Only the FIRST call per process + does anything; the rest return immediately. + """ + global _bg_started + log = logger or _LOGGER + if not autodownload_enabled(): + return + if find_managed_binary() is not None: + return # already installed -- nothing to do, stay quiet + with _bg_lock: + if _bg_started: + return + _bg_started = True + + def _run(): + try: + log.info("OpenCode not installed — installing now (%s)...", reason or "first use") + path = download_managed_binary() + if path: + log.info("OpenCode installed: %s", path) + else: + log.info( + "OpenCode install could not complete (offline?); it will be " + "retried automatically the next time the agent is used." + ) + except Exception as exc: # pragma: no cover - best-effort + log.debug("OpenCode background install failed: %s", exc) + + threading.Thread(target=_run, name="opencode-install", daemon=True).start() diff --git a/weightslab/opencode_process.py b/weightslab/opencode_process.py index d36c89b5..92db1eee 100644 --- a/weightslab/opencode_process.py +++ b/weightslab/opencode_process.py @@ -43,6 +43,8 @@ from pathlib import Path from typing import Optional +from weightslab import opencode_binary + _LOGGER = logging.getLogger(__name__) # Generous: a cold `npx` run downloads the package before the server binds. @@ -68,6 +70,17 @@ # or where a specific port is the one that happens to be forwarded/published. PORT_ENV_VAR = "WEIGHTSLAB_OPENCODE_PORT" +# Host the spawned OpenCode server BINDS to. Loopback by default -- the server +# has filesystem access and must never be reachable off the machine on a normal +# local run. But in a container reached over an SSH tunnel / published port, the +# browser's request arrives on the container's network interface, not its +# loopback, so a 127.0.0.1-only bind is refused. Setting this to 0.0.0.0 (done +# in the weightslab dev container) lets the published port reach it. Only the +# BIND host changes; the URL handed to the browser stays 127.0.0.1 (which the +# tunnel maps), so this never widens what address the page is told to use. +HOST_ENV_VAR = "WEIGHTSLAB_OPENCODE_HOST" +DEFAULT_OPENCODE_HOST = "127.0.0.1" + # Dropped directly in the workspace directory, next to (and alongside) the # AGENTS.md the landing-page agent already seeds there -- same "lives with # the experiment" reasoning, and it means deleting/moving the experiment @@ -234,6 +247,17 @@ def write_lock(workspace_dir: str, url: str, pid: Optional[int] = None) -> None: _LOGGER.warning("Could not write OpenCode lock file under %s", workspace_dir) +def opencode_bind_host() -> str: + """Host the spawned OpenCode server binds to (``--hostname``). + + ``WEIGHTSLAB_OPENCODE_HOST`` overrides the loopback default -- set it to + ``0.0.0.0`` so a container's published port / an SSH tunnel can reach the + server. Distinct from the URL reported to the browser, which stays + ``127.0.0.1`` on purpose (see HOST_ENV_VAR). + """ + return os.environ.get(HOST_ENV_VAR, "").strip() or DEFAULT_OPENCODE_HOST + + def default_opencode_port() -> int: """The port a fresh spawn asks for first -- DEFAULT_OPENCODE_PORT unless WEIGHTSLAB_OPENCODE_PORT overrides it. A malformed or out-of-range value is @@ -302,15 +326,28 @@ def pick_opencode_port() -> int: def resolve_opencode_argv() -> Optional[list]: - """Locate a way to run OpenCode, preferring an already-installed binary. - - Falls back to ``npx --yes``, which fetches the package into the npx - cache on first use -- deliberately not a global ``npm i -g``, which can - need elevated permissions and mutates the user's toolchain silently. + """Locate a way to run OpenCode, in preference order. + + 1. A weightslab-managed binary already provisioned on this machine + (``opencode_binary``) -- what makes a clean ``pip install weightslab`` + work with no Node and no manual OpenCode install. + 2. An ``opencode`` the user already has on PATH (a global/dev install): + respected before we spend bandwidth provisioning our own. + 3. Provisioning the managed binary now (a one-time ~180 MB fetch from the + npm registry, no Node required). + 4. ``npx --yes`` as a last resort -- fetches into the npx cache on first + use; deliberately not a global ``npm i -g`` (needs elevated perms and + mutates the user's toolchain silently). Requires Node. """ + managed = opencode_binary.find_managed_binary() + if managed: + return [str(managed)] exe = shutil.which("opencode") if exe: return [exe] + managed = opencode_binary.ensure_managed_binary() + if managed: + return [str(managed)] npx = shutil.which("npx") if npx: return [npx, "--yes", "opencode-ai@latest"] @@ -399,8 +436,10 @@ def resolve_or_spawn_opencode(workspace_dir: str, origin: Optional[str] = None, if argv is None: return { "ok": False, - "error": "Could not find `opencode` or `npx`. Install Node.js 20+ " - "(which provides npx), or `npm i -g opencode-ai`.", + "error": "Could not provision OpenCode: the managed binary download " + "failed (offline?) and no `opencode`/`npx` was found. Restore " + "network access, or install Node.js 20+ (provides npx), or " + "`npm i -g opencode-ai`.", } port = pick_opencode_port() @@ -408,7 +447,7 @@ def resolve_or_spawn_opencode(workspace_dir: str, origin: Optional[str] = None, for value in DEFAULT_CORS_ORIGINS: if value not in cors: cors.append(value) - cmd = argv + ["serve", "--hostname", "127.0.0.1", "--port", str(port)] + cmd = argv + ["serve", "--hostname", opencode_bind_host(), "--port", str(port)] for value in cors: cmd += ["--cors", value] diff --git a/weightslab/ui/server.py b/weightslab/ui/server.py index c6ff40af..53c7eec7 100644 --- a/weightslab/ui/server.py +++ b/weightslab/ui/server.py @@ -80,6 +80,50 @@ _LOOPBACK_ADDRESSES = {"127.0.0.1", "::1"} +def _parse_trusted_client_nets() -> list: + """Extra source networks allowed to hit the local-only control routes, + from WEIGHTSLAB_UI_TRUSTED_HOSTS (comma-separated IPs or CIDRs). + + Default is empty: loopback stays the only trusted source. This exists for + the container-behind-a-tunnel case -- when the browser reaches the UI via a + published port, the request arrives from the container's gateway, not + 127.0.0.1, so those routes would 403. There the real trust boundary is the + SSH tunnel + the host publishing only to 127.0.0.1, so trusting the internal + docker network (e.g. "172.16.0.0/12") is safe and must be opted into + explicitly. The weightslab dev container sets this. + """ + import ipaddress + raw = os.environ.get("WEIGHTSLAB_UI_TRUSTED_HOSTS", "") + nets = [] + for token in raw.split(","): + token = token.strip() + if not token: + continue + try: + nets.append(ipaddress.ip_network(token, strict=False)) + except ValueError: + logger.warning("Ignoring invalid WEIGHTSLAB_UI_TRUSTED_HOSTS entry: %r", token) + return nets + + +_TRUSTED_CLIENT_NETS = _parse_trusted_client_nets() + + +def _client_is_trusted(addr: str) -> bool: + """True if a request from ``addr`` may hit the local-only control routes: + always for loopback, and for any network in WEIGHTSLAB_UI_TRUSTED_HOSTS.""" + if addr in _LOOPBACK_ADDRESSES: + return True + if not _TRUSTED_CLIENT_NETS: + return False + import ipaddress + try: + ip = ipaddress.ip_address(addr) + except ValueError: + return False + return any(ip in net for net in _TRUSTED_CLIENT_NETS) + + def static_dir() -> str: """Absolute path to the bundled SPA directory (``weightslab/ui/static``).""" return os.path.join(os.path.dirname(os.path.abspath(__file__)), "static") @@ -470,20 +514,16 @@ def get(self) -> dict: def _resolve_opencode_argv() -> Optional[list]: - """Locate a way to run OpenCode, preferring an already-installed binary. - - Falls back to ``npx --yes``, which fetches the package into the npx cache on - first use. That is deliberately *not* ``npm install -g``: a global install may - need elevated permissions and mutates the user's toolchain behind their back, - while the npx path needs neither and is equally automatic. + """Locate a way to run OpenCode. + + Delegates to ``opencode_process.resolve_opencode_argv`` so this UI server and + the backend SDK agent share ONE resolution order: a weightslab-managed binary + (provisioned on demand for a Node-free ``pip install``) first, then an + ``opencode`` already on PATH, then a managed provision, then the ``npx --yes`` + fallback. Keeping the two paths identical is what stops them disagreeing about + which OpenCode to run for the same workspace. """ - exe = shutil.which("opencode") - if exe: - return [exe] - npx = shutil.which("npx") - if npx: - return [npx, "--yes", "opencode-ai@latest"] - return None + return opencode_process.resolve_opencode_argv() def _opencode_healthy(base_url: str, timeout: float = 1.5) -> bool: @@ -611,13 +651,16 @@ def ensure(self, workspace_dir: str, origin: Optional[str]) -> dict: argv = _resolve_opencode_argv() if argv is None: self._error = ( - "Could not find `opencode` or `npx`. Install Node.js 20+ " - "(which provides npx), or `npm i -g opencode-ai`." + "Could not provision OpenCode: the managed binary download " + "failed (offline?) and no `opencode`/`npx` was found. Restore " + "network access, or install Node.js 20+ (provides npx), or " + "`npm i -g opencode-ai`." ) return {"ok": False, "error": self._error} port = _pick_opencode_port() - cmd = argv + ["serve", "--hostname", "127.0.0.1", "--port", str(port)] + cmd = argv + ["serve", "--hostname", opencode_process.opencode_bind_host(), + "--port", str(port)] for value in _cors_origin_variants(origin): cmd += ["--cors", value] @@ -1740,7 +1783,7 @@ def _start_agent_server(self): like every other local-machine action in this server -- this one starts a process with filesystem access, so it must never be reachable off-host. """ - if self.client_address[0] not in _LOOPBACK_ADDRESSES: + if not _client_is_trusted(self.client_address[0]): self._send_json(HTTPStatus.FORBIDDEN, {"ok": False, "error": "Only reachable from localhost."}) return @@ -1770,7 +1813,7 @@ def _start_loop(self): command in the connected-experiment agent bar). Loopback-only, same reasoning as _start_agent_server -- this also starts/reuses that same process.""" - if self.client_address[0] not in _LOOPBACK_ADDRESSES: + if not _client_is_trusted(self.client_address[0]): self._send_json(HTTPStatus.FORBIDDEN, {"ok": False, "error": "Only reachable from localhost."}) return @@ -1803,7 +1846,7 @@ def _start_loop(self): self._send_json(status, result) def _stop_loop(self, loop_id: str): - if self.client_address[0] not in _LOOPBACK_ADDRESSES: + if not _client_is_trusted(self.client_address[0]): self._send_json(HTTPStatus.FORBIDDEN, {"ok": False, "error": "Only reachable from localhost."}) return @@ -1814,7 +1857,7 @@ def _stop_loop(self, loop_id: str): def _update_loop(self, loop_id: str): """Change a running loop's prompt and/or interval (the panel's Edit action) without stopping and restarting the job.""" - if self.client_address[0] not in _LOOPBACK_ADDRESSES: + if not _client_is_trusted(self.client_address[0]): self._send_json(HTTPStatus.FORBIDDEN, {"ok": False, "error": "Only reachable from localhost."}) return @@ -1837,7 +1880,7 @@ def _get_loop_messages(self, loop_id: str): """Backs a loop tab's read-only transcript: its scrollback is just this job's own OpenCode session history, written to solely by the scheduled check-in (_fire) -- nothing to merge here, only fetching.""" - if self.client_address[0] not in _LOOPBACK_ADDRESSES: + if not _client_is_trusted(self.client_address[0]): self._send_json(HTTPStatus.FORBIDDEN, {"ok": False, "error": "Only reachable from localhost."}) return @@ -1867,7 +1910,7 @@ def _data_query(self): protobuf body to forward, and this one starts from a plain JSON {query, accumulate} instead. """ - if self.client_address[0] not in _LOOPBACK_ADDRESSES: + if not _client_is_trusted(self.client_address[0]): self._send_json(HTTPStatus.FORBIDDEN, {"ok": False, "error": "Only reachable from localhost."}) return @@ -1921,7 +1964,7 @@ def _get_latest_data_query(self): _data_query/_latest_data_query). {"seq": 0} if nothing has run yet this process; the frontend only reacts when seq is NEWER than the last one it already handled.""" - if self.client_address[0] not in _LOOPBACK_ADDRESSES: + if not _client_is_trusted(self.client_address[0]): self._send_json(HTTPStatus.FORBIDDEN, {"ok": False, "error": "Only reachable from localhost."}) return @@ -1933,7 +1976,7 @@ def _track_process(self): _TrackedProcesses' own docstring for why a detached process needs this instead of being reachable through the normal process-tree kill every OTHER child of this server already gets.""" - if self.client_address[0] not in _LOOPBACK_ADDRESSES: + if not _client_is_trusted(self.client_address[0]): self._send_json(HTTPStatus.FORBIDDEN, {"ok": False, "error": "Only reachable from localhost."}) return @@ -1978,7 +2021,7 @@ def _start_local_notebook(self): # browser can't do itself; this is the one piece of local-only # control surface that requires it. Loopback-gated like the other # "local machine" actions in this server. - if self.client_address[0] not in _LOOPBACK_ADDRESSES: + if not _client_is_trusted(self.client_address[0]): self._send_json(HTTPStatus.FORBIDDEN, {"ok": False, "error": "Only reachable from localhost."}) return From 454d039da52729fe620a7736913154f1dc0ea9bd Mon Sep 17 00:00:00 2001 From: Guillaume Date: Wed, 26 Aug 2026 12:26:08 +0200 Subject: [PATCH 02/29] Fix/v2.1 UI fixes (#300) * fix(v2.1): video-gen report media + server-authoritative subview flag - reporting: add a "Generated Media" section (poster thumbnails per media field) so video/image-generation runs show their artifacts instead of an empty report; guard media columns in the Distributions path (no longer mislabelled "no numeric values"); surface a swallowed get_combined_df error as a warning so a broken dataframe isn't silently hidden. - data_service/proto: DataSamplesResponse gains is_subview/view_count/ total_count, stamped from the backend's _is_filtered state on every GetDataSamples, so a fresh client can render the subview warning ribbon with no cached UI state. Regenerated pb2 with grpcio-tools 1.68 (gencode 5.28.1). - tests: report media/distribution-guard coverage. Co-Authored-By: Claude Opus 4.8 (1M context) * fix(v2.1): backend min/max curve decimation so spikes survive 10k->1k get_signal_history_downsampled emitted the earliest-step point per bucket, so a spike between bucket edges (unless separately flagged as marker/note/outlier) was dropped server-side before the browser ever saw it. Now it emits each bucket's min-value AND max-value rows (min/max decimation); the bucket count is halved so the total stays ~max_points. 10k points -> ~1k, with spikes. + test: a non-flagged mid-bucket value spike survives, output stays ~max_points. * Check OS and adapt cmd --- tests/backend/test_logger_scale.py | 40 +++ tests/test_opencode_process.py | 10 + tests/test_reporting_media.py | 65 +++++ weightslab/backend/logger.py | 40 ++- weightslab/proto/experiment_service.proto | 8 + weightslab/proto/experiment_service_pb2.py | 262 ++++++++++---------- weightslab/reporting.py | 146 +++++++++++ weightslab/src.py | 8 +- weightslab/trainer/services/data_service.py | 30 ++- 9 files changed, 469 insertions(+), 140 deletions(-) create mode 100644 tests/test_reporting_media.py diff --git a/tests/backend/test_logger_scale.py b/tests/backend/test_logger_scale.py index d67b8d4c..5a6f8d2f 100644 --- a/tests/backend/test_logger_scale.py +++ b/tests/backend/test_logger_scale.py @@ -345,6 +345,46 @@ def test_special_points_survive_decimation(big_logger): "outlier-bearing steps were decimated away" +def test_value_spike_survives_decimation(tmp_path): + """A tall spike that is NOT flagged (no marker/note/outlier) and does not sit + on a bucket edge must still survive — min/max decimation keeps each bucket's + extreme, whereas earliest-per-bucket dropped it. Also: 10k -> ~max_points.""" + db = tmp_path / "spike.duckdb" + lg = LoggerQueue(register=False, db_path=str(db)) + lg.chkpt_manager = None + n, spike_step, spike_val = 10000, 3737, 999.0 + lg._conn.execute( + f""" + INSERT INTO signals ( + metric_name, experiment_hash, step, metric_value, timestamp, + audit_mode, is_evaluation_marker, split_name, evaluation_tags, + point_note, outliers, outlier_count, sample_count, + trend_value, trend_margin, value_min, value_max, seq) + SELECT 'loss', 'run0', t.i::INTEGER, + CASE WHEN t.i = {spike_step} THEN {spike_val} + ELSE 1.0 + 0.01 * sin(t.i / 9.0) END, + 1787000000 + t.i, FALSE, FALSE, 'train', '[]', + '', '', 0, 32, NULL, NULL, NULL, NULL, t.i + FROM range(0, {n}) AS t(i) + """ + ) + try: + hist = lg.get_signal_history_downsampled(max_points=1000) + entries = [e for per_hash in hist.values() + for steps in per_hash.values() + for lst in steps.values() for e in lst] + values = [e.get("metric_value") for e in entries] + assert any(abs((v or 0) - spike_val) < 1e-6 for v in values), \ + "value spike was decimated away" + # ~max_points, not the full 10k and not a handful. + assert 200 <= len(entries) <= 2200, f"unexpected emitted count {len(entries)}" + finally: + try: + lg.stop_background_flush() + except Exception: + pass + + def test_entries_carry_full_metadata(big_logger): """Each rendered point must arrive with the metadata the UI draws with.""" metric, h = next(iter(big_logger.truth)) diff --git a/tests/test_opencode_process.py b/tests/test_opencode_process.py index 9bc89e6f..f854b6ce 100644 --- a/tests/test_opencode_process.py +++ b/tests/test_opencode_process.py @@ -16,6 +16,7 @@ import json import os import signal +import subprocess import sys import tempfile import unittest @@ -47,6 +48,15 @@ def stop_workspace_server(workspace_dir): return if not pid: return + if os.name == "nt": + try: + subprocess.run( + ["taskkill", "/T", "/F", "/PID", str(pid)], + stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, + ) + except Exception: + pass + return try: os.killpg(os.getpgid(pid), signal.SIGTERM) except OSError: diff --git a/tests/test_reporting_media.py b/tests/test_reporting_media.py new file mode 100644 index 00000000..e551c1ce --- /dev/null +++ b/tests/test_reporting_media.py @@ -0,0 +1,65 @@ +"""Report generation for media (video/image) use cases: the report must surface +generated media instead of ignoring it, and must not mislabel a media column as +an empty numeric signal in the Distributions section.""" +import io +import unittest + +import pandas as pd + +from weightslab import reporting +from weightslab.data import media_store + + +def _png(color=(200, 60, 60)): + from PIL import Image + b = io.BytesIO() + Image.new("RGB", (16, 16), color).save(b, "PNG") + return b.getvalue() + + +class MediaReportTests(unittest.TestCase): + def setUp(self): + media_store.clear() + self.addCleanup(media_store.clear) + for sid in range(3): + media_store.put("pred_video", sid, data=b"FAKE", mime="video/mp4", + kind="video", poster=_png()) + self.df = pd.DataFrame( + {"media:pred_video": [media_store.descriptor_json(media_store.get("pred_video", s)) + for s in range(3)], + "signals//train/fm_loss": [0.5, 0.3, 0.9]}, + index=pd.Index([0, 1, 2], name="sample_id"), + ) + + def test_compute_media_examples_finds_field_and_posters(self): + examples = reporting.compute_media_examples(self.df) + self.assertEqual(len(examples), 1) + ex = examples[0] + self.assertEqual(ex["field"], "pred_video") + self.assertEqual(ex["kind"], "video") + self.assertEqual(ex["count"], 3) + self.assertTrue(ex["thumbnails"], "expected poster thumbnails") + self.assertTrue(ex["thumbnails"][0]["poster_uri"].startswith("data:image/")) + + def test_media_examples_empty_without_media(self): + plain = pd.DataFrame({"signals//loss": [1.0, 2.0]}) + self.assertEqual(reporting.compute_media_examples(plain), []) + + def test_media_section_html_renders_and_is_empty_when_none(self): + html = reporting._media_section_html(reporting.compute_media_examples(self.df)) + self.assertIn("Generated Media", html) + self.assertIn("pred_video", html) + self.assertIn("data:image/", html) + self.assertEqual(reporting._media_section_html([]), "") + + def test_distribution_on_media_column_is_flagged_not_empty_numeric(self): + entries = reporting.compute_distribution_entries(self.df, ["pred_video"], plt=None) + self.assertEqual(len(entries), 1) + self.assertTrue(entries[0].get("is_media")) + card = reporting._distribution_card_html(entries[0], "b0") + self.assertIn("is a media column", card) + self.assertNotIn("No numeric values logged", card) + + +if __name__ == "__main__": + unittest.main() diff --git a/weightslab/backend/logger.py b/weightslab/backend/logger.py index 979715d4..7ce7febe 100644 --- a/weightslab/backend/logger.py +++ b/weightslab/backend/logger.py @@ -1487,10 +1487,13 @@ def get_signal_history_downsampled( actually drawable -- not by the table size. This is what makes a hundred-million-row history openable. - The curve is split into ``max_points`` equal step-buckets and one - representative (the bucket's earliest step) is emitted per bucket, via a - streaming hash aggregate rather than a sort/window. Three further rules - keep the reduced curve faithful: + The curve is split into ``max_points / 2`` equal step-buckets and TWO + representatives — the bucket's minimum-value and maximum-value rows — are + emitted per bucket (min/max decimation), via a streaming hash aggregate + rather than a sort/window. Keeping both extremes is what preserves spikes + that fall between bucket edges (earliest-per-bucket used to drop them); + halving the bucket count keeps the total at ~``max_points`` per curve. + Three further rules keep the reduced curve faithful: * the curve's true first and last steps are always emitted, so endpoints and the x-extent never move under downsampling; @@ -1511,14 +1514,19 @@ def get_signal_history_downsampled( keep_special: emit marker/annotated/outlier rows regardless of bucketing. """ - n_buckets = max(int(max_points or _DEFAULT_MAX_POINTS_PER_CURVE), + # Two representatives per bucket (the min-value and max-value rows, see + # `reps` below) preserve spikes, so halve the bucket count to still net + # ~max_points per curve overall. 10k points -> 500 buckets x {min,max} + # -> ~1000 points, WITH the spikes that plain earliest-per-bucket + # decimation used to drop. + n_buckets = max(int(max_points or _DEFAULT_MAX_POINTS_PER_CURVE) // 2, _MIN_POINTS_PER_CURVE) where, params = self._scope_filters(metric_names, exp_hashes, x_min, x_max) cols = ", ".join(_SIGNAL_READ_COLS) # arg_min(col, step) picks each column from the bucket's earliest-step # row, so a representative is one real row rather than a blend. Rows # sharing a step within a bucket tie arbitrarily -- acceptable for a - # decimated view, and the zoom path resolves them. + # decimated view, and the zoom path resolves them. Used for the endpoints. picks = ", ".join( f"arg_min({c}, step) AS {c}" for c in _SIGNAL_READ_COLS if c not in ("metric_name", "experiment_hash", "step") @@ -1527,6 +1535,18 @@ def get_signal_history_downsampled( f"arg_max({c}, step) AS {c}" for c in _SIGNAL_READ_COLS if c not in ("metric_name", "experiment_hash", "step") ) + # Per-bucket VALUE extremes: the whole row at the bucket's min metric_value + # and the whole row at its max metric_value (step included, via + # arg_min/arg_max over metric_value). This is min/max decimation — a spike + # is by definition its bucket's extreme, so it always survives. + picks_vmin = ", ".join( + f"arg_min({c}, metric_value) AS {c}" for c in _SIGNAL_READ_COLS + if c not in ("metric_name", "experiment_hash") + ) + picks_vmax = ", ".join( + f"arg_max({c}, metric_value) AS {c}" for c in _SIGNAL_READ_COLS + if c not in ("metric_name", "experiment_hash") + ) sql = f""" WITH scoped AS ( SELECT {cols} FROM signals WHERE 1=1{where} @@ -1560,7 +1580,13 @@ def get_signal_history_downsampled( ON s.metric_name = b.m AND s.experiment_hash IS NOT DISTINCT FROM b.h ), reps AS ( - SELECT metric_name, experiment_hash, MIN(step) AS step, {picks} + -- Min/max decimation: keep BOTH the lowest- and highest-value row of + -- each bucket so up- and down-spikes between bucket edges survive + -- (plain earliest-per-bucket dropped them). + SELECT metric_name, experiment_hash, {picks_vmin} + FROM tagged GROUP BY metric_name, experiment_hash, bucket + UNION ALL + SELECT metric_name, experiment_hash, {picks_vmax} FROM tagged GROUP BY metric_name, experiment_hash, bucket ), ends AS ( diff --git a/weightslab/proto/experiment_service.proto b/weightslab/proto/experiment_service.proto index 337c51d4..17386eb8 100644 --- a/weightslab/proto/experiment_service.proto +++ b/weightslab/proto/experiment_service.proto @@ -539,6 +539,14 @@ message DataSamplesResponse { bool success = 1; string message = 2; repeated DataRecord data_records = 3; + // True when the server is currently serving a FILTERED/agent-generated + // subview of the dataset rather than the full dataset. Server-authoritative + // so a fresh client (private window, no cached UI state) can still show the + // "you are viewing a subview" warning ribbon and offer a reset. Mirrors the + // backend's own _is_filtered flag. + bool is_subview = 4; + int64 view_count = 5; // rows in the current (possibly filtered) view + int64 total_count = 6; // rows in the full dataset } // --- Server-side histogram binning --- diff --git a/weightslab/proto/experiment_service_pb2.py b/weightslab/proto/experiment_service_pb2.py index 5030a4df..3709cc25 100644 --- a/weightslab/proto/experiment_service_pb2.py +++ b/weightslab/proto/experiment_service_pb2.py @@ -24,7 +24,7 @@ -DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(b'\n)weightslab/proto/experiment_service.proto\"\x81\x02\n\x1aGetLatestLoggerDataRequest\x12\x1c\n\x14request_full_history\x18\x01 \x01(\x08\x12\x12\n\nmax_points\x18\x02 \x01(\x05\x12\x17\n\x0f\x62reak_by_slices\x18\x03 \x01(\x08\x12\x0c\n\x04tags\x18\x04 \x03(\t\x12\x12\n\ngraph_name\x18\x05 \x01(\t\x12\r\n\x05x_min\x18\x06 \x01(\x03\x12\r\n\x05x_max\x18\x07 \x01(\x03\x12\x13\n\x0bhas_x_range\x18\x08 \x01(\x08\x12\x14\n\x0cmetric_names\x18\t \x03(\t\x12\x19\n\x11\x65xperiment_hashes\x18\n \x03(\t\x12\x12\n\nindex_only\x18\x0b \x01(\x08\"\xa2\x01\n\x10SignalCurveIndex\x12\x13\n\x0bmetric_name\x18\x01 \x01(\t\x12\x17\n\x0f\x65xperiment_hash\x18\x02 \x01(\t\x12\x12\n\nfirst_step\x18\x03 \x01(\x03\x12\x11\n\tlast_step\x18\x04 \x01(\x03\x12\x13\n\x0bpoint_count\x18\x05 \x01(\x03\x12\x11\n\tvalue_min\x18\x06 \x01(\x01\x12\x11\n\tvalue_max\x18\x07 \x01(\x01\"1\n\rSignalOutlier\x12\x11\n\tsample_id\x18\x01 \x01(\t\x12\r\n\x05value\x18\x02 \x01(\x02\"\xd2\x03\n\x0fLoggerDataPoint\x12\x13\n\x0bmetric_name\x18\x01 \x01(\t\x12\x11\n\tmodel_age\x18\x02 \x01(\x05\x12\x14\n\x0cmetric_value\x18\x03 \x01(\x02\x12\x17\n\x0f\x65xperiment_hash\x18\x04 \x01(\t\x12\x11\n\ttimestamp\x18\x05 \x01(\x03\x12\x11\n\tsample_id\x18\x06 \x01(\t\x12\x1c\n\x14is_evaluation_marker\x18\x07 \x01(\x08\x12\x12\n\nsplit_name\x18\x08 \x01(\t\x12\x17\n\x0f\x65valuation_tags\x18\t \x03(\t\x12\x12\n\npoint_note\x18\n \x01(\t\x12\x12\n\naudit_mode\x18\x0b \x01(\x08\x12 \n\x08outliers\x18\x0c \x03(\x0b\x32\x0e.SignalOutlier\x12\x15\n\routlier_count\x18\r \x01(\x05\x12\x14\n\x0csample_count\x18\x0e \x01(\x05\x12\x13\n\x0btrend_value\x18\x0f \x01(\x02\x12\x14\n\x0ctrend_margin\x18\x10 \x01(\x02\x12\x16\n\x0ehas_trend_band\x18\x11 \x01(\x08\x12\x11\n\tvalue_min\x18\x12 \x01(\x02\x12\x11\n\tvalue_max\x18\x13 \x01(\x02\x12\x17\n\x0fhas_value_range\x18\x14 \x01(\x08\"\x9a\x01\n\x1bGetLatestLoggerDataResponse\x12 \n\x06points\x18\x01 \x03(\x0b\x32\x10.LoggerDataPoint\x12\x1a\n\x12weightslab_version\x18\x02 \x01(\t\x12!\n\x06\x63urves\x18\x03 \x03(\x0b\x32\x11.SignalCurveIndex\x12\x1a\n\x12\x61pplied_max_points\x18\x04 \x01(\x05\"\x07\n\x05\x45mpty\"/\n\x08NeuronId\x12\x10\n\x08layer_id\x18\x01 \x01(\x05\x12\x11\n\tneuron_id\x18\x02 \x01(\x05\"\x91\x02\n\x0fWeightOperation\x12*\n\x07op_type\x18\x01 \x01(\x0e\x32\x14.WeightOperationTypeH\x00\x88\x01\x01\x12\x15\n\x08layer_id\x18\x02 \x01(\x05H\x01\x88\x01\x01\x12\x1d\n\nneuron_ids\x18\x03 \x03(\x0b\x32\t.NeuronId\x12\x16\n\x0eneurons_to_add\x18\t \x01(\x05\x12 \n\x18zerofy_from_incoming_ids\x18\x0b \x03(\x05\x12\x1c\n\x14zerofy_to_neuron_ids\x18\x0c \x03(\x05\x12+\n\x11zerofy_predicates\x18\r \x03(\x0e\x32\x10.ZerofyPredicateB\n\n\x08_op_typeB\x0b\n\t_layer_id\"_\n\x17WeightsOperationRequest\x12/\n\x10weight_operation\x18\x01 \x01(\x0b\x32\x10.WeightOperationH\x00\x88\x01\x01\x42\x13\n\x11_weight_operation\"<\n\x18WeightsOperationResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\"\xc1\x05\n\x0fHyperParameters\x12\x1c\n\x0f\x65xperiment_name\x18\x01 \x01(\tH\x00\x88\x01\x01\x12!\n\x14training_steps_to_do\x18\x02 \x01(\x05H\x01\x88\x01\x01\x12\x1a\n\rlearning_rate\x18\x03 \x01(\x02H\x02\x88\x01\x01\x12\x17\n\nbatch_size\x18\x04 \x01(\x05H\x03\x88\x01\x01\x12 \n\x13\x66ull_eval_frequency\x18\x05 \x01(\x05H\x04\x88\x01\x01\x12 \n\x13\x63heckpont_frequency\x18\x06 \x01(\x05H\x05\x88\x01\x01\x12\x18\n\x0bis_training\x18\x07 \x01(\x08H\x06\x88\x01\x01\x12\x15\n\x08nb_steps\x18\x08 \x01(\x05H\x07\x88\x01\x01\x12\x19\n\x0c\x61uditor_mode\x18\t \x01(\x08H\x08\x88\x01\x01\x12\x1d\n\x10train_batch_size\x18\n \x01(\x05H\t\x88\x01\x01\x12\x1b\n\x0eval_batch_size\x18\x0b \x01(\x05H\n\x88\x01\x01\x12\x1c\n\x0ftest_batch_size\x18\x0c \x01(\x05H\x0b\x88\x01\x01\x12\x1c\n\x0f\x65valuation_mode\x18\r \x01(\x08H\x0c\x88\x01\x01\x12\x1e\n\x11\x65valuation_config\x18\x0e \x01(\tH\r\x88\x01\x01\x42\x12\n\x10_experiment_nameB\x17\n\x15_training_steps_to_doB\x10\n\x0e_learning_rateB\r\n\x0b_batch_sizeB\x16\n\x14_full_eval_frequencyB\x16\n\x14_checkpont_frequencyB\x0e\n\x0c_is_trainingB\x0b\n\t_nb_stepsB\x0f\n\r_auditor_modeB\x13\n\x11_train_batch_sizeB\x11\n\x0f_val_batch_sizeB\x12\n\x10_test_batch_sizeB\x12\n\x10_evaluation_modeB\x14\n\x12_evaluation_config\",\n\rMetricsStatus\x12\x0c\n\x04name\x18\x01 \x01(\t\x12\r\n\x05value\x18\x02 \x01(\x02\"~\n\rAnnotatStatus\x12\x0c\n\x04name\x18\x01 \x01(\t\x12.\n\x08metadata\x18\x02 \x03(\x0b\x32\x1c.AnnotatStatus.MetadataEntry\x1a/\n\rMetadataEntry\x12\x0b\n\x03key\x18\x01 \x01(\t\x12\r\n\x05value\x18\x02 \x01(\x02:\x02\x38\x01\"\x90\x02\n\x10TrainingStatusEx\x12\x16\n\ttimestamp\x18\x01 \x01(\tH\x00\x88\x01\x01\x12\x1c\n\x0f\x65xperiment_name\x18\x02 \x01(\tH\x01\x88\x01\x01\x12\x16\n\tmodel_age\x18\x03 \x01(\x05H\x02\x88\x01\x01\x12+\n\x0emetrics_status\x18\x04 \x01(\x0b\x32\x0e.MetricsStatusH\x03\x88\x01\x01\x12+\n\x0e\x61nnotat_status\x18\x05 \x01(\x0b\x32\x0e.AnnotatStatusH\x04\x88\x01\x01\x42\x0c\n\n_timestampB\x12\n\x10_experiment_nameB\x0c\n\n_model_ageB\x11\n\x0f_metrics_statusB\x11\n\x0f_annotat_status\"]\n\x15HyperParameterCommand\x12/\n\x10hyper_parameters\x18\x01 \x01(\x0b\x32\x10.HyperParametersH\x00\x88\x01\x01\x42\x13\n\x11_hyper_parameters\">\n\x14\x44\x65nySamplesOperation\x12\x12\n\nsample_ids\x18\x01 \x03(\t\x12\x12\n\naccumulate\x18\x02 \x01(\x08\"0\n\x17LoadCheckpointOperation\x12\x15\n\rcheckpoint_id\x18\x01 \x01(\x05\"b\n\x11PlotNoteOperation\x12\x13\n\x0bmetric_name\x18\x01 \x01(\t\x12\x17\n\x0f\x65xperiment_hash\x18\x02 \x01(\t\x12\x11\n\tmodel_age\x18\x03 \x01(\x05\x12\x0c\n\x04note\x18\x04 \x01(\t\"L\n\x17SaveCheckpointOperation\x12\x19\n\x11save_architecture\x18\x01 \x01(\x08\x12\x16\n\x0esave_optimizer\x18\x02 \x01(\x08\"\x1a\n\x18RestartInstanceOperation\"\x8d\x08\n\x0eTrainerCommand\x12\x1c\n\x14get_hyper_parameters\x18\x04 \x01(\x08\x12\x1e\n\x16get_interactive_layers\x18\x05 \x01(\x08\x12\x1d\n\x10get_data_records\x18\x06 \x01(\tH\x00\x88\x01\x01\x12%\n\x18get_single_layer_info_id\x18\x08 \x01(\x05H\x01\x88\x01\x01\x12;\n\x16hyper_parameter_change\x18\x01 \x01(\x0b\x32\x16.HyperParameterCommandH\x02\x88\x01\x01\x12:\n\x16\x64\x65ny_samples_operation\x18\x07 \x01(\x0b\x32\x15.DenySamplesOperationH\x03\x88\x01\x01\x12?\n\x1b\x64\x65ny_eval_samples_operation\x18\n \x01(\x0b\x32\x15.DenySamplesOperationH\x04\x88\x01\x01\x12@\n\x19load_checkpoint_operation\x18\t \x01(\x0b\x32\x18.LoadCheckpointOperationH\x05\x88\x01\x01\x12\x42\n\x1eremove_from_denylist_operation\x18\x0b \x01(\x0b\x32\x15.DenySamplesOperationH\x06\x88\x01\x01\x12G\n#remove_eval_from_denylist_operation\x18\x0c \x01(\x0b\x32\x15.DenySamplesOperationH\x07\x88\x01\x01\x12\x34\n\x13plot_note_operation\x18\r \x01(\x0b\x32\x12.PlotNoteOperationH\x08\x88\x01\x01\x12@\n\x19save_checkpoint_operation\x18\x0e \x01(\x0b\x32\x18.SaveCheckpointOperationH\t\x88\x01\x01\x12\x39\n\x11restart_operation\x18\x0f \x01(\x0b\x32\x19.RestartInstanceOperationH\n\x88\x01\x01\x42\x13\n\x11_get_data_recordsB\x1b\n\x19_get_single_layer_info_idB\x19\n\x17_hyper_parameter_changeB\x19\n\x17_deny_samples_operationB\x1e\n\x1c_deny_eval_samples_operationB\x1c\n\x1a_load_checkpoint_operationB!\n\x1f_remove_from_denylist_operationB&\n$_remove_eval_from_denylist_operationB\x16\n\x14_plot_note_operationB\x1c\n\x1a_save_checkpoint_operationB\x14\n\x12_restart_operation\"\x9d\x01\n\x12HyperParameterDesc\x12\r\n\x05label\x18\x01 \x01(\t\x12\x0c\n\x04name\x18\x02 \x01(\t\x12\x0c\n\x04type\x18\x03 \x01(\t\x12\x1c\n\x0fnumerical_value\x18\x04 \x01(\x02H\x00\x88\x01\x01\x12\x19\n\x0cstring_value\x18\x05 \x01(\tH\x01\x88\x01\x01\x42\x12\n\x10_numerical_valueB\x0f\n\r_string_value\"\xf2\x02\n\x10NeuronStatistics\x12!\n\tneuron_id\x18\x01 \x01(\x0b\x32\t.NeuronIdH\x00\x88\x01\x01\x12\x17\n\nneuron_age\x18\x02 \x01(\x05H\x01\x88\x01\x01\x12\x1f\n\x12train_trigger_rate\x18\x03 \x01(\x02H\x02\x88\x01\x01\x12\x1e\n\x11\x65val_trigger_rate\x18\x04 \x01(\x02H\x03\x88\x01\x01\x12\x1a\n\rlearning_rate\x18\x07 \x01(\x02H\x04\x88\x01\x01\x12\x36\n\x0bincoming_lr\x18\x08 \x03(\x0b\x32!.NeuronStatistics.IncomingLrEntry\x1a\x31\n\x0fIncomingLrEntry\x12\x0b\n\x03key\x18\x01 \x01(\x05\x12\r\n\x05value\x18\x02 \x01(\x02:\x02\x38\x01\x42\x0c\n\n_neuron_idB\r\n\x0b_neuron_ageB\x15\n\x13_train_trigger_rateB\x14\n\x12_eval_trigger_rateB\x10\n\x0e_learning_rate\"\xf0\x02\n\x13LayerRepresentation\x12\x15\n\x08layer_id\x18\x01 \x01(\x05H\x00\x88\x01\x01\x12\x17\n\nlayer_name\x18\x02 \x01(\tH\x01\x88\x01\x01\x12\x17\n\nlayer_type\x18\x03 \x01(\tH\x02\x88\x01\x01\x12\x1a\n\rneurons_count\x18\x04 \x01(\x05H\x03\x88\x01\x01\x12#\n\x16incoming_neurons_count\x18\x05 \x01(\x05H\x04\x88\x01\x01\x12\x18\n\x0bkernel_size\x18\x06 \x01(\x05H\x05\x88\x01\x01\x12\x13\n\x06stride\x18\x07 \x01(\x05H\x06\x88\x01\x01\x12-\n\x12neurons_statistics\x18\n \x03(\x0b\x32\x11.NeuronStatisticsB\x0b\n\t_layer_idB\r\n\x0b_layer_nameB\r\n\x0b_layer_typeB\x10\n\x0e_neurons_countB\x19\n\x17_incoming_neurons_countB\x0e\n\x0c_kernel_sizeB\t\n\x07_stride\"H\n\x11\x41\x63tivationRequest\x12\x10\n\x08layer_id\x18\x01 \x01(\x05\x12\x11\n\tsample_id\x18\x02 \x01(\t\x12\x0e\n\x06origin\x18\x03 \x01(\t\"H\n\rActivationMap\x12\x11\n\tneuron_id\x18\x01 \x01(\x05\x12\x0e\n\x06values\x18\x02 \x03(\x02\x12\t\n\x01H\x18\x03 \x01(\x05\x12\t\n\x01W\x18\x04 \x01(\x05\"d\n\x12\x41\x63tivationResponse\x12\x12\n\nlayer_type\x18\x01 \x01(\t\x12\x15\n\rneurons_count\x18\x02 \x01(\x05\x12#\n\x0b\x61\x63tivations\x18\x03 \x03(\x0b\x32\x0e.ActivationMap\"\x93\x01\n\tTaskField\x12\x0c\n\x04name\x18\x01 \x01(\t\x12\x15\n\x0b\x66loat_value\x18\x02 \x01(\x02H\x00\x12\x13\n\tint_value\x18\x03 \x01(\x05H\x00\x12\x16\n\x0cstring_value\x18\x04 \x01(\tH\x00\x12\x15\n\x0b\x62ytes_value\x18\x05 \x01(\x0cH\x00\x12\x14\n\nbool_value\x18\x06 \x01(\x08H\x00\x42\x07\n\x05value\"\x87\x03\n\x0eRecordMetadata\x12\x11\n\tsample_id\x18\x01 \x01(\t\x12\x14\n\x0csample_label\x18\x02 \x03(\x05\x12\x19\n\x11sample_prediction\x18\x03 \x03(\x05\x12=\n\x10sample_last_loss\x18\x04 \x03(\x0b\x32#.RecordMetadata.SampleLastLossEntry\x12\x19\n\x11sample_encounters\x18\x05 \x01(\x05\x12\x18\n\x10sample_discarded\x18\x06 \x01(\x08\x12 \n\x0c\x65xtra_fields\x18\x07 \x03(\x0b\x32\n.TaskField\x12\x16\n\x0eprediction_raw\x18\t \x01(\x0c\x12\x11\n\ttask_type\x18\n \x01(\t\x12\x19\n\x11sample_label_text\x18\x0b \x03(\t\x12\x1e\n\x16sample_prediction_text\x18\x0c \x03(\t\x1a\x35\n\x13SampleLastLossEntry\x12\x0b\n\x03key\x18\x01 \x01(\t\x12\r\n\x05value\x18\x02 \x01(\x02:\x02\x38\x01\"\x93\x01\n\x10SampleStatistics\x12\x13\n\x06origin\x18\x06 \x01(\tH\x00\x88\x01\x01\x12\x19\n\x0csample_count\x18\x07 \x01(\x05H\x01\x88\x01\x01\x12\x11\n\ttask_type\x18\t \x01(\t\x12 \n\x07records\x18\x08 \x03(\x0b\x32\x0f.RecordMetadataB\t\n\x07_originB\x0f\n\r_sample_count\"\xe6\x01\n\x0f\x43ommandResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\x12\x33\n\x16hyper_parameters_descs\x18\x03 \x03(\x0b\x32\x13.HyperParameterDesc\x12\x33\n\x15layer_representations\x18\x04 \x03(\x0b\x32\x14.LayerRepresentation\x12\x31\n\x11sample_statistics\x18\x05 \x01(\x0b\x32\x11.SampleStatisticsH\x00\x88\x01\x01\x42\x14\n\x12_sample_statistics\"U\n\rSampleRequest\x12\x16\n\tsample_id\x18\x01 \x01(\tH\x00\x88\x01\x01\x12\x13\n\x06origin\x18\x02 \x01(\tH\x01\x88\x01\x01\x42\x0c\n\n_sample_idB\t\n\x07_origin\"\xad\x02\n\x15SampleRequestResponse\x12\x16\n\tsample_id\x18\x01 \x01(\tH\x00\x88\x01\x01\x12\x13\n\x06origin\x18\x02 \x01(\tH\x01\x88\x01\x01\x12\x12\n\x05label\x18\x03 \x01(\x05H\x02\x88\x01\x01\x12\x11\n\x04\x64\x61ta\x18\x04 \x01(\x0cH\x03\x88\x01\x01\x12\x1a\n\rerror_message\x18\x05 \x01(\tH\x04\x88\x01\x01\x12\x15\n\x08raw_data\x18\x06 \x01(\x0cH\x05\x88\x01\x01\x12\x11\n\x04mask\x18\x07 \x01(\x0cH\x06\x88\x01\x01\x12\x17\n\nprediction\x18\x08 \x01(\x0cH\x07\x88\x01\x01\x42\x0c\n\n_sample_idB\t\n\x07_originB\x08\n\x06_labelB\x07\n\x05_dataB\x10\n\x0e_error_messageB\x0b\n\t_raw_dataB\x07\n\x05_maskB\r\n\x0b_prediction\"\x92\x01\n\x12\x42\x61tchSampleRequest\x12\x12\n\nsample_ids\x18\x01 \x03(\t\x12\x0e\n\x06origin\x18\x02 \x01(\t\x12\x19\n\x0cresize_width\x18\x03 \x01(\x05H\x00\x88\x01\x01\x12\x1a\n\rresize_height\x18\x04 \x01(\x05H\x01\x88\x01\x01\x42\x0f\n\r_resize_widthB\x10\n\x0e_resize_height\">\n\x13\x42\x61tchSampleResponse\x12\'\n\x07samples\x18\x01 \x03(\x0b\x32\x16.SampleRequestResponse\".\n\x0eWeightsRequest\x12\x1c\n\tneuron_id\x18\x01 \x01(\x0b\x32\t.NeuronId\"\x9d\x02\n\x0fWeightsResponse\x12\x1c\n\tneuron_id\x18\x01 \x01(\x0b\x32\t.NeuronId\x12\x17\n\nlayer_name\x18\x02 \x01(\tH\x00\x88\x01\x01\x12\x17\n\nlayer_type\x18\x03 \x01(\tH\x01\x88\x01\x01\x12\x10\n\x08incoming\x18\x04 \x01(\x05\x12\x10\n\x08outgoing\x18\x05 \x01(\x05\x12\x18\n\x0bkernel_size\x18\x06 \x01(\x05H\x02\x88\x01\x01\x12\x0f\n\x07weights\x18\x07 \x03(\x02\x12\x0f\n\x07success\x18\x0b \x01(\x08\x12\x1a\n\rerror_message\x18\x0c \x01(\tH\x03\x88\x01\x01\x42\r\n\x0b_layer_nameB\r\n\x0b_layer_typeB\x0e\n\x0c_kernel_sizeB\x10\n\x0e_error_message\"R\n\x10\x44\x61taQueryRequest\x12\r\n\x05query\x18\x01 \x01(\t\x12\x12\n\naccumulate\x18\x02 \x01(\x08\x12\x1b\n\x13is_natural_language\x18\x03 \x01(\x08\"5\n\x11\x43\x61tegoricalTagDef\x12\x0c\n\x04name\x18\x01 \x01(\t\x12\x12\n\ncategories\x18\x02 \x03(\t\"\xa9\x02\n\x11\x44\x61taQueryResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\x12\x1d\n\x15number_of_all_samples\x18\x03 \x01(\x05\x12%\n\x1dnumber_of_samples_in_the_loop\x18\x04 \x01(\x05\x12#\n\x1bnumber_of_discarded_samples\x18\x05 \x01(\x05\x12\x13\n\x0bunique_tags\x18\x06 \x03(\t\x12+\n\x11\x61gent_intent_type\x18\x07 \x01(\x0e\x32\x10.AgentIntentType\x12\x17\n\x0f\x61nalysis_result\x18\x08 \x01(\t\x12,\n\x10\x63\x61tegorical_tags\x18\t \x03(\x0b\x32\x12.CategoricalTagDef\"\xc2\x01\n\x12\x44\x61taSamplesRequest\x12\x13\n\x0bstart_index\x18\x01 \x01(\x05\x12\x13\n\x0brecords_cnt\x18\x02 \x01(\x05\x12 \n\x18include_transformed_data\x18\x03 \x01(\x08\x12\x18\n\x10include_raw_data\x18\x04 \x01(\x08\x12\x19\n\x11stats_to_retrieve\x18\x05 \x03(\t\x12\x14\n\x0cresize_width\x18\x06 \x01(\x05\x12\x15\n\rresize_height\x18\x07 \x01(\x05\"m\n\x08\x44\x61taStat\x12\x0c\n\x04name\x18\x01 \x01(\t\x12\x0c\n\x04type\x18\x02 \x01(\t\x12\r\n\x05shape\x18\x03 \x03(\x05\x12\r\n\x05value\x18\x04 \x03(\x02\x12\x14\n\x0cvalue_string\x18\x05 \x01(\t\x12\x11\n\tthumbnail\x18\x06 \x01(\x0c\">\n\nDataRecord\x12\x11\n\tsample_id\x18\x01 \x01(\t\x12\x1d\n\ndata_stats\x18\x02 \x03(\x0b\x32\t.DataStat\"Z\n\x13\x44\x61taSamplesResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\x12!\n\x0c\x64\x61ta_records\x18\x03 \x03(\x0b\x32\x0b.DataRecord\"C\n\x0fHistogramSubBar\x12\x0e\n\x06origin\x18\x01 \x01(\t\x12\x11\n\tdiscarded\x18\x02 \x01(\x08\x12\r\n\x05\x63ount\x18\x03 \x01(\x03\"h\n\x0cHistogramBin\x12\x0b\n\x03min\x18\x01 \x01(\x01\x12\x0b\n\x03max\x18\x02 \x01(\x01\x12\x0b\n\x03\x61vg\x18\x03 \x01(\x01\x12\r\n\x05\x63ount\x18\x04 \x01(\x03\x12\"\n\x08sub_bars\x18\x05 \x03(\x0b\x32\x10.HistogramSubBar\"[\n\x17\x43\x61tegoricalHistogramBar\x12\r\n\x05label\x18\x01 \x01(\t\x12\r\n\x05\x63ount\x18\x02 \x01(\x03\x12\"\n\x08sub_bars\x18\x03 \x03(\x0b\x32\x10.HistogramSubBar\"4\n\x10HistogramRequest\x12\x0e\n\x06\x63olumn\x18\x01 \x01(\t\x12\x10\n\x08max_bins\x18\x02 \x01(\x05\"\xb2\x01\n\x11HistogramResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\x12\x12\n\ntotal_rows\x18\x03 \x01(\x03\x12\x1b\n\x04\x62ins\x18\x04 \x03(\x0b\x32\r.HistogramBin\x12\x16\n\x0eis_categorical\x18\x05 \x01(\x08\x12\x32\n\x10\x63\x61tegorical_bars\x18\x06 \x03(\x0b\x32\x18.CategoricalHistogramBar\"W\n\x12GetMetaDataRequest\x12\x13\n\x0bstart_index\x18\x01 \x01(\x05\x12\x13\n\x0brecords_cnt\x18\x02 \x01(\x05\x12\x17\n\x0fmodal_sample_id\x18\x03 \x01(\t\"\x99\x01\n\x13GetMetaDataResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\x12\x1a\n\x12\x61ll_metadata_names\x18\x03 \x03(\t\x12!\n\x0cgrid_records\x18\x04 \x03(\x0b\x32\x0b.DataRecord\x12!\n\x0cmodal_record\x18\x05 \x01(\x0b\x32\x0b.DataRecord\"j\n\x12StepSamplesRequest\x12\x13\n\x0bmetric_name\x18\x01 \x01(\t\x12\x17\n\x0f\x65xperiment_hash\x18\x02 \x01(\t\x12\x11\n\tmodel_age\x18\x03 \x01(\x05\x12\x13\n\x0bmax_samples\x18\x04 \x01(\x05\"{\n\x13StepSamplesResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\x12\x12\n\nsample_ids\x18\x03 \x03(\t\x12\x17\n\x0ftotal_available\x18\x04 \x01(\x05\x12\x15\n\rsample_values\x18\x05 \x03(\x02\"Y\n\x1aGetSignalTrajectoryRequest\x12\x13\n\x0bsignal_name\x18\x01 \x01(\t\x12\x12\n\nsample_ids\x18\x02 \x03(\t\x12\x12\n\nmax_points\x18\x03 \x01(\x05\"4\n\x10SignalTrajectory\x12\x11\n\tsample_id\x18\x01 \x01(\t\x12\r\n\x05value\x18\x02 \x03(\x02\"}\n\x1bGetSignalTrajectoryResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\x12\x13\n\x0bsignal_name\x18\x03 \x01(\t\x12\'\n\x0ctrajectories\x18\x04 \x03(\x0b\x32\x11.SignalTrajectory\"Y\n\x11PointCloudRequest\x12\x11\n\tsample_id\x18\x01 \x01(\t\x12\x0e\n\x06origin\x18\x02 \x01(\t\x12\x12\n\nmax_points\x18\x03 \x01(\x05\x12\r\n\x05\x66ield\x18\x04 \x01(\t\"\xbf\x01\n\x0fPointCloudChunk\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\x12\x12\n\nnum_points\x18\x03 \x01(\x05\x12\x14\n\x0cnum_features\x18\x04 \x01(\x05\x12\x10\n\x08pc_range\x18\x05 \x03(\x02\x12\x0c\n\x04\x64\x61ta\x18\x06 \x01(\x0c\x12\x13\n\x0b\x63hunk_index\x18\x07 \x01(\x05\x12\x14\n\x0ctotal_chunks\x18\x08 \x01(\x05\x12\x15\n\rfeature_names\x18\t \x03(\t\"b\n\x0cMediaRequest\x12\x11\n\tsample_id\x18\x01 \x01(\t\x12\x0e\n\x06origin\x18\x02 \x01(\t\x12\x0c\n\x04kind\x18\x03 \x01(\t\x12\x12\n\nmax_frames\x18\x04 \x01(\x05\x12\r\n\x05\x66ield\x18\x05 \x01(\t\"\x92\x02\n\nMediaChunk\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\x12\x11\n\tmime_type\x18\x03 \x01(\t\x12\x13\n\x0b\x66rame_count\x18\x04 \x01(\x05\x12\x0b\n\x03\x66ps\x18\x05 \x01(\x02\x12\x11\n\thas_audio\x18\x06 \x01(\x08\x12\r\n\x05width\x18\x07 \x01(\x05\x12\x0e\n\x06height\x18\x08 \x01(\x05\x12\x13\n\x0btotal_bytes\x18\t \x01(\x05\x12\x18\n\x10\x64uration_seconds\x18\n \x01(\x02\x12\x13\n\x0bsample_rate\x18\x0b \x01(\x05\x12\x0c\n\x04\x64\x61ta\x18\x0c \x01(\x0c\x12\x13\n\x0b\x63hunk_index\x18\r \x01(\x05\x12\x14\n\x0ctotal_chunks\x18\x0e \x01(\x05\"\x8c\x02\n\x10\x44\x61taEditsRequest\x12\x11\n\tstat_name\x18\x01 \x01(\t\x12\x13\n\x0b\x66loat_value\x18\x02 \x01(\x02\x12\x14\n\x0cstring_value\x18\x03 \x01(\t\x12\x12\n\nbool_value\x18\x04 \x01(\x08\x12\x1d\n\x04type\x18\x05 \x01(\x0e\x32\x0f.SampleEditType\x12\x13\n\x0bsamples_ids\x18\x06 \x03(\t\x12\x16\n\x0esample_origins\x18\x07 \x03(\t\x12\x16\n\x0eis_categorical\x18\x08 \x01(\x08\x12\x12\n\ncategories\x18\t \x03(\t\x12\x15\n\rsample_values\x18\n \x03(\x02\x12\x17\n\x0f\x65xperiment_hash\x18\x0b \x01(\t\"5\n\x11\x44\x61taEditsResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\":\n\x12\x44\x61taSplitsResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x13\n\x0bsplit_names\x18\x02 \x03(\t\"9\n\x13\x41gentHealthResponse\x12\x11\n\tavailable\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\"^\n\x16InitializeAgentRequest\x12\x0f\n\x07\x61pi_key\x18\x01 \x01(\t\x12$\n\x08provider\x18\x02 \x01(\x0e\x32\x12.AgentProviderType\x12\r\n\x05model\x18\x03 \x01(\t\";\n\x17InitializeAgentResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\"(\n\x17\x43hangeAgentModelRequest\x12\r\n\x05model\x18\x01 \x01(\t\"<\n\x18\x43hangeAgentModelResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\"\x17\n\x15GetAgentModelsRequest\"J\n\x16GetAgentModelsResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0e\n\x06models\x18\x02 \x03(\t\x12\x0f\n\x07message\x18\x03 \x01(\t\"6\n\x12ResetAgentResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\"=\n\x19\x43learAgentHistoryResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\"?\n\x1b\x43ompactAgentHistoryResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\"\xe5\x01\n\x1cGetAgentContextUsageResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\x12\r\n\x05model\x18\x03 \x01(\t\x12\x16\n\x0e\x63ontext_window\x18\x04 \x01(\x03\x12\x14\n\x0cinput_tokens\x18\x05 \x01(\x03\x12\x15\n\routput_tokens\x18\x06 \x01(\x03\x12\x18\n\x10reasoning_tokens\x18\x07 \x01(\x03\x12\x19\n\x11\x63\x61\x63he_read_tokens\x18\x08 \x01(\x03\x12\x1a\n\x12\x63\x61\x63he_write_tokens\x18\t \x01(\x03\"3\n\x18RestoreCheckpointRequest\x12\x17\n\x0f\x65xperiment_hash\x18\x01 \x01(\t\"=\n\x19RestoreCheckpointResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\"\xa8\x01\n\x11\x45xperimentRunInfo\x12\x17\n\x0f\x65xperiment_hash\x18\x01 \x01(\t\x12\x17\n\x0f\x65xperiment_name\x18\x02 \x01(\t\x12\r\n\x05notes\x18\x03 \x01(\t\x12\x0f\n\x07\x63reated\x18\x04 \x01(\t\x12\x11\n\tlast_used\x18\x05 \x01(\t\x12\x1a\n\x12latest_weight_step\x18\x06 \x01(\x05\x12\x12\n\nis_current\x18\x07 \x01(\x08\"\x1b\n\x19ListExperimentRunsRequest\">\n\x1aListExperimentRunsResponse\x12 \n\x04runs\x18\x01 \x03(\x0b\x32\x12.ExperimentRunInfo\"G\n\x1aRenameExperimentRunRequest\x12\x17\n\x0f\x65xperiment_hash\x18\x01 \x01(\t\x12\x10\n\x08new_name\x18\x02 \x01(\t\"?\n\x1bRenameExperimentRunResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\"F\n\x1cSetExperimentRunNotesRequest\x12\x17\n\x0f\x65xperiment_hash\x18\x01 \x01(\t\x12\r\n\x05notes\x18\x02 \x01(\t\"A\n\x1dSetExperimentRunNotesResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\"R\n\x18TriggerEvaluationRequest\x12\x12\n\nsplit_name\x18\x01 \x01(\t\x12\x0c\n\x04tags\x18\x02 \x03(\t\x12\x14\n\x0cuse_full_set\x18\x03 \x01(\x08\"=\n\x19TriggerEvaluationResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\"\x1c\n\x1aGetEvaluationStatusRequest\"\x81\x01\n\x1bGetEvaluationStatusResponse\x12\x0e\n\x06status\x18\x01 \x01(\t\x12\x0f\n\x07\x63urrent\x18\x02 \x01(\x05\x12\r\n\x05total\x18\x03 \x01(\x05\x12\x0f\n\x07message\x18\x04 \x01(\t\x12\r\n\x05\x65rror\x18\x05 \x01(\t\x12\x12\n\nsplit_name\x18\x06 \x01(\t\")\n\x17\x43\x61ncelEvaluationRequest\x12\x0e\n\x06reason\x18\x01 \x01(\t\"<\n\x18\x43\x61ncelEvaluationResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\"7\n\x16RunNotebookCellRequest\x12\x0c\n\x04\x63ode\x18\x01 \x01(\t\x12\x0f\n\x07\x63\x65ll_id\x18\x02 \x01(\t\"2\n\x10NotebookCellDone\x12\x12\n\nexec_count\x18\x01 \x01(\x05\x12\n\n\x02ok\x18\x02 \x01(\x08\"\x1e\n\x1cInterruptNotebookCellRequest\":\n\x1dInterruptNotebookCellResponse\x12\n\n\x02ok\x18\x01 \x01(\x08\x12\r\n\x05\x65rror\x18\x02 \x01(\t\"\xbd\x01\n\x11NotebookCellChunk\x12\x0f\n\x07\x63\x65ll_id\x18\x01 \x01(\t\x12\x10\n\x06stdout\x18\x02 \x01(\tH\x00\x12\x10\n\x06stderr\x18\x03 \x01(\tH\x00\x12\x15\n\x0bresult_text\x18\x04 \x01(\tH\x00\x12\x13\n\timage_png\x18\x05 \x01(\x0cH\x00\x12\x19\n\x0f\x65rror_traceback\x18\x06 \x01(\tH\x00\x12!\n\x04\x64one\x18\x07 \x01(\x0b\x32\x11.NotebookCellDoneH\x00\x42\t\n\x07payload\"S\n\x10NotebookResponse\x12\x12\n\nipynb_json\x18\x01 \x01(\t\x12\x0f\n\x07\x65xisted\x18\x02 \x01(\x08\x12\x0c\n\x04path\x18\x03 \x01(\t\x12\x0c\n\x04name\x18\x04 \x01(\t\"7\n\x13SaveNotebookRequest\x12\x12\n\nipynb_json\x18\x01 \x01(\t\x12\x0c\n\x04name\x18\x02 \x01(\t\"M\n\x14SaveNotebookResponse\x12\n\n\x02ok\x18\x01 \x01(\x08\x12\x0c\n\x04path\x18\x02 \x01(\t\x12\r\n\x05\x65rror\x18\x03 \x01(\t\x12\x0c\n\x04name\x18\x04 \x01(\t\"C\n\x1bGenerateNotebookCodeRequest\x12\x0e\n\x06prompt\x18\x01 \x01(\t\x12\x14\n\x0c\x63ontext_code\x18\x02 \x01(\t\"\\\n\x1cGenerateNotebookCodeResponse\x12\x0c\n\x04\x63ode\x18\x01 \x01(\t\x12\x13\n\x0b\x65xplanation\x18\x02 \x01(\t\x12\n\n\x02ok\x18\x03 \x01(\x08\x12\r\n\x05\x65rror\x18\x04 \x01(\t\"~\n\x18\x45xportAnnotationsRequest\x12\'\n\x06\x66ormat\x18\x01 \x01(\x0e\x32\x17.AnnotationExportFormat\x12\x0e\n\x06origin\x18\x02 \x01(\t\x12\x1b\n\x13include_predictions\x18\x03 \x01(\x08\x12\x0c\n\x04tags\x18\x04 \x03(\t\"\x88\x01\n\x19\x45xportAnnotationsResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\x12\x0f\n\x07payload\x18\x03 \x01(\x0c\x12\x10\n\x08\x66ilename\x18\x04 \x01(\t\x12\x11\n\tmime_type\x18\x05 \x01(\t\x12\x13\n\x0bimage_count\x18\x06 \x01(\x05*d\n\x13WeightOperationType\x12\n\n\x06ZEROFY\x10\x00\x12\x10\n\x0cREINITIALIZE\x10\x01\x12\n\n\x06\x46REEZE\x10\x02\x12\x12\n\x0eREMOVE_NEURONS\x10\t\x12\x0f\n\x0b\x41\x44\x44_NEURONS\x10\n*o\n\x0fZerofyPredicate\x12\x19\n\x15ZEROFY_PREDICATE_NONE\x10\x00\x12 \n\x1cZEROFY_PREDICATE_WITH_FROZEN\x10\x01\x12\x1f\n\x1bZEROFY_PREDICATE_WITH_OLDER\x10\x02*M\n\x0f\x41gentIntentType\x12\x12\n\x0eINTENT_UNKNOWN\x10\x00\x12\x11\n\rINTENT_FILTER\x10\x01\x12\x13\n\x0fINTENT_ANALYSIS\x10\x02*I\n\x0eSampleEditType\x12\x11\n\rEDIT_OVERRIDE\x10\x00\x12\x13\n\x0f\x45\x44IT_ACCUMULATE\x10\x01\x12\x0f\n\x0b\x45\x44IT_REMOVE\x10\x02*C\n\x11\x41gentProviderType\x12\x17\n\x13PROVIDER_OPENROUTER\x10\x00\x12\x15\n\x11PROVIDER_OPENCODE\x10\x01*m\n\x16\x41nnotationExportFormat\x12\x16\n\x12\x45XPORT_FORMAT_CVAT\x10\x00\x12\x1e\n\x1a\x45XPORT_FORMAT_LABEL_STUDIO\x10\x01\x12\x1b\n\x17\x45XPORT_FORMAT_V7_DARWIN\x10\x02\x32\xfe\x12\n\x11\x45xperimentService\x12P\n\x13GetLatestLoggerData\x12\x1b.GetLatestLoggerDataRequest\x1a\x1c.GetLatestLoggerDataResponse\x12\x36\n\x11\x45xperimentCommand\x12\x0f.TrainerCommand\x1a\x10.CommandResponse\x12H\n\x11ManipulateWeights\x12\x18.WeightsOperationRequest\x1a\x19.WeightsOperationResponse\x12/\n\nGetWeights\x12\x0f.WeightsRequest\x1a\x10.WeightsResponse\x12\x39\n\x0eGetActivations\x12\x12.ActivationRequest\x1a\x13.ActivationResponse\x12\x37\n\nGetSamples\x12\x13.BatchSampleRequest\x1a\x14.BatchSampleResponse\x12\x37\n\x0e\x41pplyDataQuery\x12\x11.DataQueryRequest\x1a\x12.DataQueryResponse\x12;\n\x0eGetDataSamples\x12\x13.DataSamplesRequest\x1a\x14.DataSamplesResponse\x12\x35\n\x0cGetHistogram\x12\x11.HistogramRequest\x1a\x12.HistogramResponse\x12\x38\n\x0bGetMetaData\x12\x13.GetMetaDataRequest\x1a\x14.GetMetaDataResponse\x12P\n\x13GetSignalTrajectory\x12\x1b.GetSignalTrajectoryRequest\x1a\x1c.GetSignalTrajectoryResponse\x12;\n\x0eGetStepSamples\x12\x13.StepSamplesRequest\x1a\x14.StepSamplesResponse\x12\x37\n\rGetPointCloud\x12\x12.PointCloudRequest\x1a\x10.PointCloudChunk0\x01\x12(\n\x08GetMedia\x12\r.MediaRequest\x1a\x0b.MediaChunk0\x01\x12\x37\n\x0e\x45\x64itDataSample\x12\x11.DataEditsRequest\x1a\x12.DataEditsResponse\x12,\n\rGetDataSplits\x12\x06.Empty\x1a\x13.DataSplitsResponse\x12\x30\n\x10\x43heckAgentHealth\x12\x06.Empty\x1a\x14.AgentHealthResponse\x12\x44\n\x0fInitializeAgent\x12\x17.InitializeAgentRequest\x1a\x18.InitializeAgentResponse\x12G\n\x10\x43hangeAgentModel\x12\x18.ChangeAgentModelRequest\x1a\x19.ChangeAgentModelResponse\x12\x41\n\x0eGetAgentModels\x12\x16.GetAgentModelsRequest\x1a\x17.GetAgentModelsResponse\x12)\n\nResetAgent\x12\x06.Empty\x1a\x13.ResetAgentResponse\x12\x37\n\x11\x43learAgentHistory\x12\x06.Empty\x1a\x1a.ClearAgentHistoryResponse\x12;\n\x13\x43ompactAgentHistory\x12\x06.Empty\x1a\x1c.CompactAgentHistoryResponse\x12=\n\x14GetAgentContextUsage\x12\x06.Empty\x1a\x1d.GetAgentContextUsageResponse\x12@\n\x0fRunNotebookCell\x12\x17.RunNotebookCellRequest\x1a\x12.NotebookCellChunk0\x01\x12V\n\x15InterruptNotebookCell\x12\x1d.InterruptNotebookCellRequest\x1a\x1e.InterruptNotebookCellResponse\x12(\n\x0bGetNotebook\x12\x06.Empty\x1a\x11.NotebookResponse\x12;\n\x0cSaveNotebook\x12\x14.SaveNotebookRequest\x1a\x15.SaveNotebookResponse\x12S\n\x14GenerateNotebookCode\x12\x1c.GenerateNotebookCodeRequest\x1a\x1d.GenerateNotebookCodeResponse\x12J\n\x11RestoreCheckpoint\x12\x19.RestoreCheckpointRequest\x1a\x1a.RestoreCheckpointResponse\x12M\n\x12ListExperimentRuns\x12\x1a.ListExperimentRunsRequest\x1a\x1b.ListExperimentRunsResponse\x12P\n\x13RenameExperimentRun\x12\x1b.RenameExperimentRunRequest\x1a\x1c.RenameExperimentRunResponse\x12V\n\x15SetExperimentRunNotes\x12\x1d.SetExperimentRunNotesRequest\x1a\x1e.SetExperimentRunNotesResponse\x12J\n\x11TriggerEvaluation\x12\x19.TriggerEvaluationRequest\x1a\x1a.TriggerEvaluationResponse\x12P\n\x13GetEvaluationStatus\x12\x1b.GetEvaluationStatusRequest\x1a\x1c.GetEvaluationStatusResponse\x12G\n\x10\x43\x61ncelEvaluation\x12\x18.CancelEvaluationRequest\x1a\x19.CancelEvaluationResponse\x12J\n\x11\x45xportAnnotations\x12\x19.ExportAnnotationsRequest\x1a\x1a.ExportAnnotationsResponseb\x06proto3') +DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(b'\n)weightslab/proto/experiment_service.proto\"\x81\x02\n\x1aGetLatestLoggerDataRequest\x12\x1c\n\x14request_full_history\x18\x01 \x01(\x08\x12\x12\n\nmax_points\x18\x02 \x01(\x05\x12\x17\n\x0f\x62reak_by_slices\x18\x03 \x01(\x08\x12\x0c\n\x04tags\x18\x04 \x03(\t\x12\x12\n\ngraph_name\x18\x05 \x01(\t\x12\r\n\x05x_min\x18\x06 \x01(\x03\x12\r\n\x05x_max\x18\x07 \x01(\x03\x12\x13\n\x0bhas_x_range\x18\x08 \x01(\x08\x12\x14\n\x0cmetric_names\x18\t \x03(\t\x12\x19\n\x11\x65xperiment_hashes\x18\n \x03(\t\x12\x12\n\nindex_only\x18\x0b \x01(\x08\"\xa2\x01\n\x10SignalCurveIndex\x12\x13\n\x0bmetric_name\x18\x01 \x01(\t\x12\x17\n\x0f\x65xperiment_hash\x18\x02 \x01(\t\x12\x12\n\nfirst_step\x18\x03 \x01(\x03\x12\x11\n\tlast_step\x18\x04 \x01(\x03\x12\x13\n\x0bpoint_count\x18\x05 \x01(\x03\x12\x11\n\tvalue_min\x18\x06 \x01(\x01\x12\x11\n\tvalue_max\x18\x07 \x01(\x01\"1\n\rSignalOutlier\x12\x11\n\tsample_id\x18\x01 \x01(\t\x12\r\n\x05value\x18\x02 \x01(\x02\"\xd2\x03\n\x0fLoggerDataPoint\x12\x13\n\x0bmetric_name\x18\x01 \x01(\t\x12\x11\n\tmodel_age\x18\x02 \x01(\x05\x12\x14\n\x0cmetric_value\x18\x03 \x01(\x02\x12\x17\n\x0f\x65xperiment_hash\x18\x04 \x01(\t\x12\x11\n\ttimestamp\x18\x05 \x01(\x03\x12\x11\n\tsample_id\x18\x06 \x01(\t\x12\x1c\n\x14is_evaluation_marker\x18\x07 \x01(\x08\x12\x12\n\nsplit_name\x18\x08 \x01(\t\x12\x17\n\x0f\x65valuation_tags\x18\t \x03(\t\x12\x12\n\npoint_note\x18\n \x01(\t\x12\x12\n\naudit_mode\x18\x0b \x01(\x08\x12 \n\x08outliers\x18\x0c \x03(\x0b\x32\x0e.SignalOutlier\x12\x15\n\routlier_count\x18\r \x01(\x05\x12\x14\n\x0csample_count\x18\x0e \x01(\x05\x12\x13\n\x0btrend_value\x18\x0f \x01(\x02\x12\x14\n\x0ctrend_margin\x18\x10 \x01(\x02\x12\x16\n\x0ehas_trend_band\x18\x11 \x01(\x08\x12\x11\n\tvalue_min\x18\x12 \x01(\x02\x12\x11\n\tvalue_max\x18\x13 \x01(\x02\x12\x17\n\x0fhas_value_range\x18\x14 \x01(\x08\"\x9a\x01\n\x1bGetLatestLoggerDataResponse\x12 \n\x06points\x18\x01 \x03(\x0b\x32\x10.LoggerDataPoint\x12\x1a\n\x12weightslab_version\x18\x02 \x01(\t\x12!\n\x06\x63urves\x18\x03 \x03(\x0b\x32\x11.SignalCurveIndex\x12\x1a\n\x12\x61pplied_max_points\x18\x04 \x01(\x05\"\x07\n\x05\x45mpty\"/\n\x08NeuronId\x12\x10\n\x08layer_id\x18\x01 \x01(\x05\x12\x11\n\tneuron_id\x18\x02 \x01(\x05\"\x91\x02\n\x0fWeightOperation\x12*\n\x07op_type\x18\x01 \x01(\x0e\x32\x14.WeightOperationTypeH\x00\x88\x01\x01\x12\x15\n\x08layer_id\x18\x02 \x01(\x05H\x01\x88\x01\x01\x12\x1d\n\nneuron_ids\x18\x03 \x03(\x0b\x32\t.NeuronId\x12\x16\n\x0eneurons_to_add\x18\t \x01(\x05\x12 \n\x18zerofy_from_incoming_ids\x18\x0b \x03(\x05\x12\x1c\n\x14zerofy_to_neuron_ids\x18\x0c \x03(\x05\x12+\n\x11zerofy_predicates\x18\r \x03(\x0e\x32\x10.ZerofyPredicateB\n\n\x08_op_typeB\x0b\n\t_layer_id\"_\n\x17WeightsOperationRequest\x12/\n\x10weight_operation\x18\x01 \x01(\x0b\x32\x10.WeightOperationH\x00\x88\x01\x01\x42\x13\n\x11_weight_operation\"<\n\x18WeightsOperationResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\"\xc1\x05\n\x0fHyperParameters\x12\x1c\n\x0f\x65xperiment_name\x18\x01 \x01(\tH\x00\x88\x01\x01\x12!\n\x14training_steps_to_do\x18\x02 \x01(\x05H\x01\x88\x01\x01\x12\x1a\n\rlearning_rate\x18\x03 \x01(\x02H\x02\x88\x01\x01\x12\x17\n\nbatch_size\x18\x04 \x01(\x05H\x03\x88\x01\x01\x12 \n\x13\x66ull_eval_frequency\x18\x05 \x01(\x05H\x04\x88\x01\x01\x12 \n\x13\x63heckpont_frequency\x18\x06 \x01(\x05H\x05\x88\x01\x01\x12\x18\n\x0bis_training\x18\x07 \x01(\x08H\x06\x88\x01\x01\x12\x15\n\x08nb_steps\x18\x08 \x01(\x05H\x07\x88\x01\x01\x12\x19\n\x0c\x61uditor_mode\x18\t \x01(\x08H\x08\x88\x01\x01\x12\x1d\n\x10train_batch_size\x18\n \x01(\x05H\t\x88\x01\x01\x12\x1b\n\x0eval_batch_size\x18\x0b \x01(\x05H\n\x88\x01\x01\x12\x1c\n\x0ftest_batch_size\x18\x0c \x01(\x05H\x0b\x88\x01\x01\x12\x1c\n\x0f\x65valuation_mode\x18\r \x01(\x08H\x0c\x88\x01\x01\x12\x1e\n\x11\x65valuation_config\x18\x0e \x01(\tH\r\x88\x01\x01\x42\x12\n\x10_experiment_nameB\x17\n\x15_training_steps_to_doB\x10\n\x0e_learning_rateB\r\n\x0b_batch_sizeB\x16\n\x14_full_eval_frequencyB\x16\n\x14_checkpont_frequencyB\x0e\n\x0c_is_trainingB\x0b\n\t_nb_stepsB\x0f\n\r_auditor_modeB\x13\n\x11_train_batch_sizeB\x11\n\x0f_val_batch_sizeB\x12\n\x10_test_batch_sizeB\x12\n\x10_evaluation_modeB\x14\n\x12_evaluation_config\",\n\rMetricsStatus\x12\x0c\n\x04name\x18\x01 \x01(\t\x12\r\n\x05value\x18\x02 \x01(\x02\"~\n\rAnnotatStatus\x12\x0c\n\x04name\x18\x01 \x01(\t\x12.\n\x08metadata\x18\x02 \x03(\x0b\x32\x1c.AnnotatStatus.MetadataEntry\x1a/\n\rMetadataEntry\x12\x0b\n\x03key\x18\x01 \x01(\t\x12\r\n\x05value\x18\x02 \x01(\x02:\x02\x38\x01\"\x90\x02\n\x10TrainingStatusEx\x12\x16\n\ttimestamp\x18\x01 \x01(\tH\x00\x88\x01\x01\x12\x1c\n\x0f\x65xperiment_name\x18\x02 \x01(\tH\x01\x88\x01\x01\x12\x16\n\tmodel_age\x18\x03 \x01(\x05H\x02\x88\x01\x01\x12+\n\x0emetrics_status\x18\x04 \x01(\x0b\x32\x0e.MetricsStatusH\x03\x88\x01\x01\x12+\n\x0e\x61nnotat_status\x18\x05 \x01(\x0b\x32\x0e.AnnotatStatusH\x04\x88\x01\x01\x42\x0c\n\n_timestampB\x12\n\x10_experiment_nameB\x0c\n\n_model_ageB\x11\n\x0f_metrics_statusB\x11\n\x0f_annotat_status\"]\n\x15HyperParameterCommand\x12/\n\x10hyper_parameters\x18\x01 \x01(\x0b\x32\x10.HyperParametersH\x00\x88\x01\x01\x42\x13\n\x11_hyper_parameters\">\n\x14\x44\x65nySamplesOperation\x12\x12\n\nsample_ids\x18\x01 \x03(\t\x12\x12\n\naccumulate\x18\x02 \x01(\x08\"0\n\x17LoadCheckpointOperation\x12\x15\n\rcheckpoint_id\x18\x01 \x01(\x05\"b\n\x11PlotNoteOperation\x12\x13\n\x0bmetric_name\x18\x01 \x01(\t\x12\x17\n\x0f\x65xperiment_hash\x18\x02 \x01(\t\x12\x11\n\tmodel_age\x18\x03 \x01(\x05\x12\x0c\n\x04note\x18\x04 \x01(\t\"L\n\x17SaveCheckpointOperation\x12\x19\n\x11save_architecture\x18\x01 \x01(\x08\x12\x16\n\x0esave_optimizer\x18\x02 \x01(\x08\"\x1a\n\x18RestartInstanceOperation\"\x8d\x08\n\x0eTrainerCommand\x12\x1c\n\x14get_hyper_parameters\x18\x04 \x01(\x08\x12\x1e\n\x16get_interactive_layers\x18\x05 \x01(\x08\x12\x1d\n\x10get_data_records\x18\x06 \x01(\tH\x00\x88\x01\x01\x12%\n\x18get_single_layer_info_id\x18\x08 \x01(\x05H\x01\x88\x01\x01\x12;\n\x16hyper_parameter_change\x18\x01 \x01(\x0b\x32\x16.HyperParameterCommandH\x02\x88\x01\x01\x12:\n\x16\x64\x65ny_samples_operation\x18\x07 \x01(\x0b\x32\x15.DenySamplesOperationH\x03\x88\x01\x01\x12?\n\x1b\x64\x65ny_eval_samples_operation\x18\n \x01(\x0b\x32\x15.DenySamplesOperationH\x04\x88\x01\x01\x12@\n\x19load_checkpoint_operation\x18\t \x01(\x0b\x32\x18.LoadCheckpointOperationH\x05\x88\x01\x01\x12\x42\n\x1eremove_from_denylist_operation\x18\x0b \x01(\x0b\x32\x15.DenySamplesOperationH\x06\x88\x01\x01\x12G\n#remove_eval_from_denylist_operation\x18\x0c \x01(\x0b\x32\x15.DenySamplesOperationH\x07\x88\x01\x01\x12\x34\n\x13plot_note_operation\x18\r \x01(\x0b\x32\x12.PlotNoteOperationH\x08\x88\x01\x01\x12@\n\x19save_checkpoint_operation\x18\x0e \x01(\x0b\x32\x18.SaveCheckpointOperationH\t\x88\x01\x01\x12\x39\n\x11restart_operation\x18\x0f \x01(\x0b\x32\x19.RestartInstanceOperationH\n\x88\x01\x01\x42\x13\n\x11_get_data_recordsB\x1b\n\x19_get_single_layer_info_idB\x19\n\x17_hyper_parameter_changeB\x19\n\x17_deny_samples_operationB\x1e\n\x1c_deny_eval_samples_operationB\x1c\n\x1a_load_checkpoint_operationB!\n\x1f_remove_from_denylist_operationB&\n$_remove_eval_from_denylist_operationB\x16\n\x14_plot_note_operationB\x1c\n\x1a_save_checkpoint_operationB\x14\n\x12_restart_operation\"\x9d\x01\n\x12HyperParameterDesc\x12\r\n\x05label\x18\x01 \x01(\t\x12\x0c\n\x04name\x18\x02 \x01(\t\x12\x0c\n\x04type\x18\x03 \x01(\t\x12\x1c\n\x0fnumerical_value\x18\x04 \x01(\x02H\x00\x88\x01\x01\x12\x19\n\x0cstring_value\x18\x05 \x01(\tH\x01\x88\x01\x01\x42\x12\n\x10_numerical_valueB\x0f\n\r_string_value\"\xf2\x02\n\x10NeuronStatistics\x12!\n\tneuron_id\x18\x01 \x01(\x0b\x32\t.NeuronIdH\x00\x88\x01\x01\x12\x17\n\nneuron_age\x18\x02 \x01(\x05H\x01\x88\x01\x01\x12\x1f\n\x12train_trigger_rate\x18\x03 \x01(\x02H\x02\x88\x01\x01\x12\x1e\n\x11\x65val_trigger_rate\x18\x04 \x01(\x02H\x03\x88\x01\x01\x12\x1a\n\rlearning_rate\x18\x07 \x01(\x02H\x04\x88\x01\x01\x12\x36\n\x0bincoming_lr\x18\x08 \x03(\x0b\x32!.NeuronStatistics.IncomingLrEntry\x1a\x31\n\x0fIncomingLrEntry\x12\x0b\n\x03key\x18\x01 \x01(\x05\x12\r\n\x05value\x18\x02 \x01(\x02:\x02\x38\x01\x42\x0c\n\n_neuron_idB\r\n\x0b_neuron_ageB\x15\n\x13_train_trigger_rateB\x14\n\x12_eval_trigger_rateB\x10\n\x0e_learning_rate\"\xf0\x02\n\x13LayerRepresentation\x12\x15\n\x08layer_id\x18\x01 \x01(\x05H\x00\x88\x01\x01\x12\x17\n\nlayer_name\x18\x02 \x01(\tH\x01\x88\x01\x01\x12\x17\n\nlayer_type\x18\x03 \x01(\tH\x02\x88\x01\x01\x12\x1a\n\rneurons_count\x18\x04 \x01(\x05H\x03\x88\x01\x01\x12#\n\x16incoming_neurons_count\x18\x05 \x01(\x05H\x04\x88\x01\x01\x12\x18\n\x0bkernel_size\x18\x06 \x01(\x05H\x05\x88\x01\x01\x12\x13\n\x06stride\x18\x07 \x01(\x05H\x06\x88\x01\x01\x12-\n\x12neurons_statistics\x18\n \x03(\x0b\x32\x11.NeuronStatisticsB\x0b\n\t_layer_idB\r\n\x0b_layer_nameB\r\n\x0b_layer_typeB\x10\n\x0e_neurons_countB\x19\n\x17_incoming_neurons_countB\x0e\n\x0c_kernel_sizeB\t\n\x07_stride\"H\n\x11\x41\x63tivationRequest\x12\x10\n\x08layer_id\x18\x01 \x01(\x05\x12\x11\n\tsample_id\x18\x02 \x01(\t\x12\x0e\n\x06origin\x18\x03 \x01(\t\"H\n\rActivationMap\x12\x11\n\tneuron_id\x18\x01 \x01(\x05\x12\x0e\n\x06values\x18\x02 \x03(\x02\x12\t\n\x01H\x18\x03 \x01(\x05\x12\t\n\x01W\x18\x04 \x01(\x05\"d\n\x12\x41\x63tivationResponse\x12\x12\n\nlayer_type\x18\x01 \x01(\t\x12\x15\n\rneurons_count\x18\x02 \x01(\x05\x12#\n\x0b\x61\x63tivations\x18\x03 \x03(\x0b\x32\x0e.ActivationMap\"\x93\x01\n\tTaskField\x12\x0c\n\x04name\x18\x01 \x01(\t\x12\x15\n\x0b\x66loat_value\x18\x02 \x01(\x02H\x00\x12\x13\n\tint_value\x18\x03 \x01(\x05H\x00\x12\x16\n\x0cstring_value\x18\x04 \x01(\tH\x00\x12\x15\n\x0b\x62ytes_value\x18\x05 \x01(\x0cH\x00\x12\x14\n\nbool_value\x18\x06 \x01(\x08H\x00\x42\x07\n\x05value\"\x87\x03\n\x0eRecordMetadata\x12\x11\n\tsample_id\x18\x01 \x01(\t\x12\x14\n\x0csample_label\x18\x02 \x03(\x05\x12\x19\n\x11sample_prediction\x18\x03 \x03(\x05\x12=\n\x10sample_last_loss\x18\x04 \x03(\x0b\x32#.RecordMetadata.SampleLastLossEntry\x12\x19\n\x11sample_encounters\x18\x05 \x01(\x05\x12\x18\n\x10sample_discarded\x18\x06 \x01(\x08\x12 \n\x0c\x65xtra_fields\x18\x07 \x03(\x0b\x32\n.TaskField\x12\x16\n\x0eprediction_raw\x18\t \x01(\x0c\x12\x11\n\ttask_type\x18\n \x01(\t\x12\x19\n\x11sample_label_text\x18\x0b \x03(\t\x12\x1e\n\x16sample_prediction_text\x18\x0c \x03(\t\x1a\x35\n\x13SampleLastLossEntry\x12\x0b\n\x03key\x18\x01 \x01(\t\x12\r\n\x05value\x18\x02 \x01(\x02:\x02\x38\x01\"\x93\x01\n\x10SampleStatistics\x12\x13\n\x06origin\x18\x06 \x01(\tH\x00\x88\x01\x01\x12\x19\n\x0csample_count\x18\x07 \x01(\x05H\x01\x88\x01\x01\x12\x11\n\ttask_type\x18\t \x01(\t\x12 \n\x07records\x18\x08 \x03(\x0b\x32\x0f.RecordMetadataB\t\n\x07_originB\x0f\n\r_sample_count\"\xe6\x01\n\x0f\x43ommandResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\x12\x33\n\x16hyper_parameters_descs\x18\x03 \x03(\x0b\x32\x13.HyperParameterDesc\x12\x33\n\x15layer_representations\x18\x04 \x03(\x0b\x32\x14.LayerRepresentation\x12\x31\n\x11sample_statistics\x18\x05 \x01(\x0b\x32\x11.SampleStatisticsH\x00\x88\x01\x01\x42\x14\n\x12_sample_statistics\"U\n\rSampleRequest\x12\x16\n\tsample_id\x18\x01 \x01(\tH\x00\x88\x01\x01\x12\x13\n\x06origin\x18\x02 \x01(\tH\x01\x88\x01\x01\x42\x0c\n\n_sample_idB\t\n\x07_origin\"\xad\x02\n\x15SampleRequestResponse\x12\x16\n\tsample_id\x18\x01 \x01(\tH\x00\x88\x01\x01\x12\x13\n\x06origin\x18\x02 \x01(\tH\x01\x88\x01\x01\x12\x12\n\x05label\x18\x03 \x01(\x05H\x02\x88\x01\x01\x12\x11\n\x04\x64\x61ta\x18\x04 \x01(\x0cH\x03\x88\x01\x01\x12\x1a\n\rerror_message\x18\x05 \x01(\tH\x04\x88\x01\x01\x12\x15\n\x08raw_data\x18\x06 \x01(\x0cH\x05\x88\x01\x01\x12\x11\n\x04mask\x18\x07 \x01(\x0cH\x06\x88\x01\x01\x12\x17\n\nprediction\x18\x08 \x01(\x0cH\x07\x88\x01\x01\x42\x0c\n\n_sample_idB\t\n\x07_originB\x08\n\x06_labelB\x07\n\x05_dataB\x10\n\x0e_error_messageB\x0b\n\t_raw_dataB\x07\n\x05_maskB\r\n\x0b_prediction\"\x92\x01\n\x12\x42\x61tchSampleRequest\x12\x12\n\nsample_ids\x18\x01 \x03(\t\x12\x0e\n\x06origin\x18\x02 \x01(\t\x12\x19\n\x0cresize_width\x18\x03 \x01(\x05H\x00\x88\x01\x01\x12\x1a\n\rresize_height\x18\x04 \x01(\x05H\x01\x88\x01\x01\x42\x0f\n\r_resize_widthB\x10\n\x0e_resize_height\">\n\x13\x42\x61tchSampleResponse\x12\'\n\x07samples\x18\x01 \x03(\x0b\x32\x16.SampleRequestResponse\".\n\x0eWeightsRequest\x12\x1c\n\tneuron_id\x18\x01 \x01(\x0b\x32\t.NeuronId\"\x9d\x02\n\x0fWeightsResponse\x12\x1c\n\tneuron_id\x18\x01 \x01(\x0b\x32\t.NeuronId\x12\x17\n\nlayer_name\x18\x02 \x01(\tH\x00\x88\x01\x01\x12\x17\n\nlayer_type\x18\x03 \x01(\tH\x01\x88\x01\x01\x12\x10\n\x08incoming\x18\x04 \x01(\x05\x12\x10\n\x08outgoing\x18\x05 \x01(\x05\x12\x18\n\x0bkernel_size\x18\x06 \x01(\x05H\x02\x88\x01\x01\x12\x0f\n\x07weights\x18\x07 \x03(\x02\x12\x0f\n\x07success\x18\x0b \x01(\x08\x12\x1a\n\rerror_message\x18\x0c \x01(\tH\x03\x88\x01\x01\x42\r\n\x0b_layer_nameB\r\n\x0b_layer_typeB\x0e\n\x0c_kernel_sizeB\x10\n\x0e_error_message\"R\n\x10\x44\x61taQueryRequest\x12\r\n\x05query\x18\x01 \x01(\t\x12\x12\n\naccumulate\x18\x02 \x01(\x08\x12\x1b\n\x13is_natural_language\x18\x03 \x01(\x08\"5\n\x11\x43\x61tegoricalTagDef\x12\x0c\n\x04name\x18\x01 \x01(\t\x12\x12\n\ncategories\x18\x02 \x03(\t\"\xa9\x02\n\x11\x44\x61taQueryResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\x12\x1d\n\x15number_of_all_samples\x18\x03 \x01(\x05\x12%\n\x1dnumber_of_samples_in_the_loop\x18\x04 \x01(\x05\x12#\n\x1bnumber_of_discarded_samples\x18\x05 \x01(\x05\x12\x13\n\x0bunique_tags\x18\x06 \x03(\t\x12+\n\x11\x61gent_intent_type\x18\x07 \x01(\x0e\x32\x10.AgentIntentType\x12\x17\n\x0f\x61nalysis_result\x18\x08 \x01(\t\x12,\n\x10\x63\x61tegorical_tags\x18\t \x03(\x0b\x32\x12.CategoricalTagDef\"\xc2\x01\n\x12\x44\x61taSamplesRequest\x12\x13\n\x0bstart_index\x18\x01 \x01(\x05\x12\x13\n\x0brecords_cnt\x18\x02 \x01(\x05\x12 \n\x18include_transformed_data\x18\x03 \x01(\x08\x12\x18\n\x10include_raw_data\x18\x04 \x01(\x08\x12\x19\n\x11stats_to_retrieve\x18\x05 \x03(\t\x12\x14\n\x0cresize_width\x18\x06 \x01(\x05\x12\x15\n\rresize_height\x18\x07 \x01(\x05\"m\n\x08\x44\x61taStat\x12\x0c\n\x04name\x18\x01 \x01(\t\x12\x0c\n\x04type\x18\x02 \x01(\t\x12\r\n\x05shape\x18\x03 \x03(\x05\x12\r\n\x05value\x18\x04 \x03(\x02\x12\x14\n\x0cvalue_string\x18\x05 \x01(\t\x12\x11\n\tthumbnail\x18\x06 \x01(\x0c\">\n\nDataRecord\x12\x11\n\tsample_id\x18\x01 \x01(\t\x12\x1d\n\ndata_stats\x18\x02 \x03(\x0b\x32\t.DataStat\"\x97\x01\n\x13\x44\x61taSamplesResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\x12!\n\x0c\x64\x61ta_records\x18\x03 \x03(\x0b\x32\x0b.DataRecord\x12\x12\n\nis_subview\x18\x04 \x01(\x08\x12\x12\n\nview_count\x18\x05 \x01(\x03\x12\x13\n\x0btotal_count\x18\x06 \x01(\x03\"C\n\x0fHistogramSubBar\x12\x0e\n\x06origin\x18\x01 \x01(\t\x12\x11\n\tdiscarded\x18\x02 \x01(\x08\x12\r\n\x05\x63ount\x18\x03 \x01(\x03\"h\n\x0cHistogramBin\x12\x0b\n\x03min\x18\x01 \x01(\x01\x12\x0b\n\x03max\x18\x02 \x01(\x01\x12\x0b\n\x03\x61vg\x18\x03 \x01(\x01\x12\r\n\x05\x63ount\x18\x04 \x01(\x03\x12\"\n\x08sub_bars\x18\x05 \x03(\x0b\x32\x10.HistogramSubBar\"[\n\x17\x43\x61tegoricalHistogramBar\x12\r\n\x05label\x18\x01 \x01(\t\x12\r\n\x05\x63ount\x18\x02 \x01(\x03\x12\"\n\x08sub_bars\x18\x03 \x03(\x0b\x32\x10.HistogramSubBar\"4\n\x10HistogramRequest\x12\x0e\n\x06\x63olumn\x18\x01 \x01(\t\x12\x10\n\x08max_bins\x18\x02 \x01(\x05\"\xb2\x01\n\x11HistogramResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\x12\x12\n\ntotal_rows\x18\x03 \x01(\x03\x12\x1b\n\x04\x62ins\x18\x04 \x03(\x0b\x32\r.HistogramBin\x12\x16\n\x0eis_categorical\x18\x05 \x01(\x08\x12\x32\n\x10\x63\x61tegorical_bars\x18\x06 \x03(\x0b\x32\x18.CategoricalHistogramBar\"W\n\x12GetMetaDataRequest\x12\x13\n\x0bstart_index\x18\x01 \x01(\x05\x12\x13\n\x0brecords_cnt\x18\x02 \x01(\x05\x12\x17\n\x0fmodal_sample_id\x18\x03 \x01(\t\"\x99\x01\n\x13GetMetaDataResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\x12\x1a\n\x12\x61ll_metadata_names\x18\x03 \x03(\t\x12!\n\x0cgrid_records\x18\x04 \x03(\x0b\x32\x0b.DataRecord\x12!\n\x0cmodal_record\x18\x05 \x01(\x0b\x32\x0b.DataRecord\"j\n\x12StepSamplesRequest\x12\x13\n\x0bmetric_name\x18\x01 \x01(\t\x12\x17\n\x0f\x65xperiment_hash\x18\x02 \x01(\t\x12\x11\n\tmodel_age\x18\x03 \x01(\x05\x12\x13\n\x0bmax_samples\x18\x04 \x01(\x05\"{\n\x13StepSamplesResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\x12\x12\n\nsample_ids\x18\x03 \x03(\t\x12\x17\n\x0ftotal_available\x18\x04 \x01(\x05\x12\x15\n\rsample_values\x18\x05 \x03(\x02\"Y\n\x1aGetSignalTrajectoryRequest\x12\x13\n\x0bsignal_name\x18\x01 \x01(\t\x12\x12\n\nsample_ids\x18\x02 \x03(\t\x12\x12\n\nmax_points\x18\x03 \x01(\x05\"4\n\x10SignalTrajectory\x12\x11\n\tsample_id\x18\x01 \x01(\t\x12\r\n\x05value\x18\x02 \x03(\x02\"}\n\x1bGetSignalTrajectoryResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\x12\x13\n\x0bsignal_name\x18\x03 \x01(\t\x12\'\n\x0ctrajectories\x18\x04 \x03(\x0b\x32\x11.SignalTrajectory\"Y\n\x11PointCloudRequest\x12\x11\n\tsample_id\x18\x01 \x01(\t\x12\x0e\n\x06origin\x18\x02 \x01(\t\x12\x12\n\nmax_points\x18\x03 \x01(\x05\x12\r\n\x05\x66ield\x18\x04 \x01(\t\"\xbf\x01\n\x0fPointCloudChunk\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\x12\x12\n\nnum_points\x18\x03 \x01(\x05\x12\x14\n\x0cnum_features\x18\x04 \x01(\x05\x12\x10\n\x08pc_range\x18\x05 \x03(\x02\x12\x0c\n\x04\x64\x61ta\x18\x06 \x01(\x0c\x12\x13\n\x0b\x63hunk_index\x18\x07 \x01(\x05\x12\x14\n\x0ctotal_chunks\x18\x08 \x01(\x05\x12\x15\n\rfeature_names\x18\t \x03(\t\"b\n\x0cMediaRequest\x12\x11\n\tsample_id\x18\x01 \x01(\t\x12\x0e\n\x06origin\x18\x02 \x01(\t\x12\x0c\n\x04kind\x18\x03 \x01(\t\x12\x12\n\nmax_frames\x18\x04 \x01(\x05\x12\r\n\x05\x66ield\x18\x05 \x01(\t\"\x92\x02\n\nMediaChunk\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\x12\x11\n\tmime_type\x18\x03 \x01(\t\x12\x13\n\x0b\x66rame_count\x18\x04 \x01(\x05\x12\x0b\n\x03\x66ps\x18\x05 \x01(\x02\x12\x11\n\thas_audio\x18\x06 \x01(\x08\x12\r\n\x05width\x18\x07 \x01(\x05\x12\x0e\n\x06height\x18\x08 \x01(\x05\x12\x13\n\x0btotal_bytes\x18\t \x01(\x05\x12\x18\n\x10\x64uration_seconds\x18\n \x01(\x02\x12\x13\n\x0bsample_rate\x18\x0b \x01(\x05\x12\x0c\n\x04\x64\x61ta\x18\x0c \x01(\x0c\x12\x13\n\x0b\x63hunk_index\x18\r \x01(\x05\x12\x14\n\x0ctotal_chunks\x18\x0e \x01(\x05\"\x8c\x02\n\x10\x44\x61taEditsRequest\x12\x11\n\tstat_name\x18\x01 \x01(\t\x12\x13\n\x0b\x66loat_value\x18\x02 \x01(\x02\x12\x14\n\x0cstring_value\x18\x03 \x01(\t\x12\x12\n\nbool_value\x18\x04 \x01(\x08\x12\x1d\n\x04type\x18\x05 \x01(\x0e\x32\x0f.SampleEditType\x12\x13\n\x0bsamples_ids\x18\x06 \x03(\t\x12\x16\n\x0esample_origins\x18\x07 \x03(\t\x12\x16\n\x0eis_categorical\x18\x08 \x01(\x08\x12\x12\n\ncategories\x18\t \x03(\t\x12\x15\n\rsample_values\x18\n \x03(\x02\x12\x17\n\x0f\x65xperiment_hash\x18\x0b \x01(\t\"5\n\x11\x44\x61taEditsResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\":\n\x12\x44\x61taSplitsResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x13\n\x0bsplit_names\x18\x02 \x03(\t\"9\n\x13\x41gentHealthResponse\x12\x11\n\tavailable\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\"^\n\x16InitializeAgentRequest\x12\x0f\n\x07\x61pi_key\x18\x01 \x01(\t\x12$\n\x08provider\x18\x02 \x01(\x0e\x32\x12.AgentProviderType\x12\r\n\x05model\x18\x03 \x01(\t\";\n\x17InitializeAgentResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\"(\n\x17\x43hangeAgentModelRequest\x12\r\n\x05model\x18\x01 \x01(\t\"<\n\x18\x43hangeAgentModelResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\"\x17\n\x15GetAgentModelsRequest\"J\n\x16GetAgentModelsResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0e\n\x06models\x18\x02 \x03(\t\x12\x0f\n\x07message\x18\x03 \x01(\t\"6\n\x12ResetAgentResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\"=\n\x19\x43learAgentHistoryResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\"?\n\x1b\x43ompactAgentHistoryResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\"\xe5\x01\n\x1cGetAgentContextUsageResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\x12\r\n\x05model\x18\x03 \x01(\t\x12\x16\n\x0e\x63ontext_window\x18\x04 \x01(\x03\x12\x14\n\x0cinput_tokens\x18\x05 \x01(\x03\x12\x15\n\routput_tokens\x18\x06 \x01(\x03\x12\x18\n\x10reasoning_tokens\x18\x07 \x01(\x03\x12\x19\n\x11\x63\x61\x63he_read_tokens\x18\x08 \x01(\x03\x12\x1a\n\x12\x63\x61\x63he_write_tokens\x18\t \x01(\x03\"3\n\x18RestoreCheckpointRequest\x12\x17\n\x0f\x65xperiment_hash\x18\x01 \x01(\t\"=\n\x19RestoreCheckpointResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\"\xa8\x01\n\x11\x45xperimentRunInfo\x12\x17\n\x0f\x65xperiment_hash\x18\x01 \x01(\t\x12\x17\n\x0f\x65xperiment_name\x18\x02 \x01(\t\x12\r\n\x05notes\x18\x03 \x01(\t\x12\x0f\n\x07\x63reated\x18\x04 \x01(\t\x12\x11\n\tlast_used\x18\x05 \x01(\t\x12\x1a\n\x12latest_weight_step\x18\x06 \x01(\x05\x12\x12\n\nis_current\x18\x07 \x01(\x08\"\x1b\n\x19ListExperimentRunsRequest\">\n\x1aListExperimentRunsResponse\x12 \n\x04runs\x18\x01 \x03(\x0b\x32\x12.ExperimentRunInfo\"G\n\x1aRenameExperimentRunRequest\x12\x17\n\x0f\x65xperiment_hash\x18\x01 \x01(\t\x12\x10\n\x08new_name\x18\x02 \x01(\t\"?\n\x1bRenameExperimentRunResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\"F\n\x1cSetExperimentRunNotesRequest\x12\x17\n\x0f\x65xperiment_hash\x18\x01 \x01(\t\x12\r\n\x05notes\x18\x02 \x01(\t\"A\n\x1dSetExperimentRunNotesResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\"R\n\x18TriggerEvaluationRequest\x12\x12\n\nsplit_name\x18\x01 \x01(\t\x12\x0c\n\x04tags\x18\x02 \x03(\t\x12\x14\n\x0cuse_full_set\x18\x03 \x01(\x08\"=\n\x19TriggerEvaluationResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\"\x1c\n\x1aGetEvaluationStatusRequest\"\x81\x01\n\x1bGetEvaluationStatusResponse\x12\x0e\n\x06status\x18\x01 \x01(\t\x12\x0f\n\x07\x63urrent\x18\x02 \x01(\x05\x12\r\n\x05total\x18\x03 \x01(\x05\x12\x0f\n\x07message\x18\x04 \x01(\t\x12\r\n\x05\x65rror\x18\x05 \x01(\t\x12\x12\n\nsplit_name\x18\x06 \x01(\t\")\n\x17\x43\x61ncelEvaluationRequest\x12\x0e\n\x06reason\x18\x01 \x01(\t\"<\n\x18\x43\x61ncelEvaluationResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\"7\n\x16RunNotebookCellRequest\x12\x0c\n\x04\x63ode\x18\x01 \x01(\t\x12\x0f\n\x07\x63\x65ll_id\x18\x02 \x01(\t\"2\n\x10NotebookCellDone\x12\x12\n\nexec_count\x18\x01 \x01(\x05\x12\n\n\x02ok\x18\x02 \x01(\x08\"\x1e\n\x1cInterruptNotebookCellRequest\":\n\x1dInterruptNotebookCellResponse\x12\n\n\x02ok\x18\x01 \x01(\x08\x12\r\n\x05\x65rror\x18\x02 \x01(\t\"\xbd\x01\n\x11NotebookCellChunk\x12\x0f\n\x07\x63\x65ll_id\x18\x01 \x01(\t\x12\x10\n\x06stdout\x18\x02 \x01(\tH\x00\x12\x10\n\x06stderr\x18\x03 \x01(\tH\x00\x12\x15\n\x0bresult_text\x18\x04 \x01(\tH\x00\x12\x13\n\timage_png\x18\x05 \x01(\x0cH\x00\x12\x19\n\x0f\x65rror_traceback\x18\x06 \x01(\tH\x00\x12!\n\x04\x64one\x18\x07 \x01(\x0b\x32\x11.NotebookCellDoneH\x00\x42\t\n\x07payload\"S\n\x10NotebookResponse\x12\x12\n\nipynb_json\x18\x01 \x01(\t\x12\x0f\n\x07\x65xisted\x18\x02 \x01(\x08\x12\x0c\n\x04path\x18\x03 \x01(\t\x12\x0c\n\x04name\x18\x04 \x01(\t\"7\n\x13SaveNotebookRequest\x12\x12\n\nipynb_json\x18\x01 \x01(\t\x12\x0c\n\x04name\x18\x02 \x01(\t\"M\n\x14SaveNotebookResponse\x12\n\n\x02ok\x18\x01 \x01(\x08\x12\x0c\n\x04path\x18\x02 \x01(\t\x12\r\n\x05\x65rror\x18\x03 \x01(\t\x12\x0c\n\x04name\x18\x04 \x01(\t\"C\n\x1bGenerateNotebookCodeRequest\x12\x0e\n\x06prompt\x18\x01 \x01(\t\x12\x14\n\x0c\x63ontext_code\x18\x02 \x01(\t\"\\\n\x1cGenerateNotebookCodeResponse\x12\x0c\n\x04\x63ode\x18\x01 \x01(\t\x12\x13\n\x0b\x65xplanation\x18\x02 \x01(\t\x12\n\n\x02ok\x18\x03 \x01(\x08\x12\r\n\x05\x65rror\x18\x04 \x01(\t\"~\n\x18\x45xportAnnotationsRequest\x12\'\n\x06\x66ormat\x18\x01 \x01(\x0e\x32\x17.AnnotationExportFormat\x12\x0e\n\x06origin\x18\x02 \x01(\t\x12\x1b\n\x13include_predictions\x18\x03 \x01(\x08\x12\x0c\n\x04tags\x18\x04 \x03(\t\"\x88\x01\n\x19\x45xportAnnotationsResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\x12\x0f\n\x07payload\x18\x03 \x01(\x0c\x12\x10\n\x08\x66ilename\x18\x04 \x01(\t\x12\x11\n\tmime_type\x18\x05 \x01(\t\x12\x13\n\x0bimage_count\x18\x06 \x01(\x05*d\n\x13WeightOperationType\x12\n\n\x06ZEROFY\x10\x00\x12\x10\n\x0cREINITIALIZE\x10\x01\x12\n\n\x06\x46REEZE\x10\x02\x12\x12\n\x0eREMOVE_NEURONS\x10\t\x12\x0f\n\x0b\x41\x44\x44_NEURONS\x10\n*o\n\x0fZerofyPredicate\x12\x19\n\x15ZEROFY_PREDICATE_NONE\x10\x00\x12 \n\x1cZEROFY_PREDICATE_WITH_FROZEN\x10\x01\x12\x1f\n\x1bZEROFY_PREDICATE_WITH_OLDER\x10\x02*M\n\x0f\x41gentIntentType\x12\x12\n\x0eINTENT_UNKNOWN\x10\x00\x12\x11\n\rINTENT_FILTER\x10\x01\x12\x13\n\x0fINTENT_ANALYSIS\x10\x02*I\n\x0eSampleEditType\x12\x11\n\rEDIT_OVERRIDE\x10\x00\x12\x13\n\x0f\x45\x44IT_ACCUMULATE\x10\x01\x12\x0f\n\x0b\x45\x44IT_REMOVE\x10\x02*C\n\x11\x41gentProviderType\x12\x17\n\x13PROVIDER_OPENROUTER\x10\x00\x12\x15\n\x11PROVIDER_OPENCODE\x10\x01*m\n\x16\x41nnotationExportFormat\x12\x16\n\x12\x45XPORT_FORMAT_CVAT\x10\x00\x12\x1e\n\x1a\x45XPORT_FORMAT_LABEL_STUDIO\x10\x01\x12\x1b\n\x17\x45XPORT_FORMAT_V7_DARWIN\x10\x02\x32\xfe\x12\n\x11\x45xperimentService\x12P\n\x13GetLatestLoggerData\x12\x1b.GetLatestLoggerDataRequest\x1a\x1c.GetLatestLoggerDataResponse\x12\x36\n\x11\x45xperimentCommand\x12\x0f.TrainerCommand\x1a\x10.CommandResponse\x12H\n\x11ManipulateWeights\x12\x18.WeightsOperationRequest\x1a\x19.WeightsOperationResponse\x12/\n\nGetWeights\x12\x0f.WeightsRequest\x1a\x10.WeightsResponse\x12\x39\n\x0eGetActivations\x12\x12.ActivationRequest\x1a\x13.ActivationResponse\x12\x37\n\nGetSamples\x12\x13.BatchSampleRequest\x1a\x14.BatchSampleResponse\x12\x37\n\x0e\x41pplyDataQuery\x12\x11.DataQueryRequest\x1a\x12.DataQueryResponse\x12;\n\x0eGetDataSamples\x12\x13.DataSamplesRequest\x1a\x14.DataSamplesResponse\x12\x35\n\x0cGetHistogram\x12\x11.HistogramRequest\x1a\x12.HistogramResponse\x12\x38\n\x0bGetMetaData\x12\x13.GetMetaDataRequest\x1a\x14.GetMetaDataResponse\x12P\n\x13GetSignalTrajectory\x12\x1b.GetSignalTrajectoryRequest\x1a\x1c.GetSignalTrajectoryResponse\x12;\n\x0eGetStepSamples\x12\x13.StepSamplesRequest\x1a\x14.StepSamplesResponse\x12\x37\n\rGetPointCloud\x12\x12.PointCloudRequest\x1a\x10.PointCloudChunk0\x01\x12(\n\x08GetMedia\x12\r.MediaRequest\x1a\x0b.MediaChunk0\x01\x12\x37\n\x0e\x45\x64itDataSample\x12\x11.DataEditsRequest\x1a\x12.DataEditsResponse\x12,\n\rGetDataSplits\x12\x06.Empty\x1a\x13.DataSplitsResponse\x12\x30\n\x10\x43heckAgentHealth\x12\x06.Empty\x1a\x14.AgentHealthResponse\x12\x44\n\x0fInitializeAgent\x12\x17.InitializeAgentRequest\x1a\x18.InitializeAgentResponse\x12G\n\x10\x43hangeAgentModel\x12\x18.ChangeAgentModelRequest\x1a\x19.ChangeAgentModelResponse\x12\x41\n\x0eGetAgentModels\x12\x16.GetAgentModelsRequest\x1a\x17.GetAgentModelsResponse\x12)\n\nResetAgent\x12\x06.Empty\x1a\x13.ResetAgentResponse\x12\x37\n\x11\x43learAgentHistory\x12\x06.Empty\x1a\x1a.ClearAgentHistoryResponse\x12;\n\x13\x43ompactAgentHistory\x12\x06.Empty\x1a\x1c.CompactAgentHistoryResponse\x12=\n\x14GetAgentContextUsage\x12\x06.Empty\x1a\x1d.GetAgentContextUsageResponse\x12@\n\x0fRunNotebookCell\x12\x17.RunNotebookCellRequest\x1a\x12.NotebookCellChunk0\x01\x12V\n\x15InterruptNotebookCell\x12\x1d.InterruptNotebookCellRequest\x1a\x1e.InterruptNotebookCellResponse\x12(\n\x0bGetNotebook\x12\x06.Empty\x1a\x11.NotebookResponse\x12;\n\x0cSaveNotebook\x12\x14.SaveNotebookRequest\x1a\x15.SaveNotebookResponse\x12S\n\x14GenerateNotebookCode\x12\x1c.GenerateNotebookCodeRequest\x1a\x1d.GenerateNotebookCodeResponse\x12J\n\x11RestoreCheckpoint\x12\x19.RestoreCheckpointRequest\x1a\x1a.RestoreCheckpointResponse\x12M\n\x12ListExperimentRuns\x12\x1a.ListExperimentRunsRequest\x1a\x1b.ListExperimentRunsResponse\x12P\n\x13RenameExperimentRun\x12\x1b.RenameExperimentRunRequest\x1a\x1c.RenameExperimentRunResponse\x12V\n\x15SetExperimentRunNotes\x12\x1d.SetExperimentRunNotesRequest\x1a\x1e.SetExperimentRunNotesResponse\x12J\n\x11TriggerEvaluation\x12\x19.TriggerEvaluationRequest\x1a\x1a.TriggerEvaluationResponse\x12P\n\x13GetEvaluationStatus\x12\x1b.GetEvaluationStatusRequest\x1a\x1c.GetEvaluationStatusResponse\x12G\n\x10\x43\x61ncelEvaluation\x12\x18.CancelEvaluationRequest\x1a\x19.CancelEvaluationResponse\x12J\n\x11\x45xportAnnotations\x12\x19.ExportAnnotationsRequest\x1a\x1a.ExportAnnotationsResponseb\x06proto3') _globals = globals() _builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals) @@ -37,18 +37,18 @@ _globals['_NEURONSTATISTICS_INCOMINGLRENTRY']._serialized_options = b'8\001' _globals['_RECORDMETADATA_SAMPLELASTLOSSENTRY']._loaded_options = None _globals['_RECORDMETADATA_SAMPLELASTLOSSENTRY']._serialized_options = b'8\001' - _globals['_WEIGHTOPERATIONTYPE']._serialized_start=13429 - _globals['_WEIGHTOPERATIONTYPE']._serialized_end=13529 - _globals['_ZEROFYPREDICATE']._serialized_start=13531 - _globals['_ZEROFYPREDICATE']._serialized_end=13642 - _globals['_AGENTINTENTTYPE']._serialized_start=13644 - _globals['_AGENTINTENTTYPE']._serialized_end=13721 - _globals['_SAMPLEEDITTYPE']._serialized_start=13723 - _globals['_SAMPLEEDITTYPE']._serialized_end=13796 - _globals['_AGENTPROVIDERTYPE']._serialized_start=13798 - _globals['_AGENTPROVIDERTYPE']._serialized_end=13865 - _globals['_ANNOTATIONEXPORTFORMAT']._serialized_start=13867 - _globals['_ANNOTATIONEXPORTFORMAT']._serialized_end=13976 + _globals['_WEIGHTOPERATIONTYPE']._serialized_start=13491 + _globals['_WEIGHTOPERATIONTYPE']._serialized_end=13591 + _globals['_ZEROFYPREDICATE']._serialized_start=13593 + _globals['_ZEROFYPREDICATE']._serialized_end=13704 + _globals['_AGENTINTENTTYPE']._serialized_start=13706 + _globals['_AGENTINTENTTYPE']._serialized_end=13783 + _globals['_SAMPLEEDITTYPE']._serialized_start=13785 + _globals['_SAMPLEEDITTYPE']._serialized_end=13858 + _globals['_AGENTPROVIDERTYPE']._serialized_start=13860 + _globals['_AGENTPROVIDERTYPE']._serialized_end=13927 + _globals['_ANNOTATIONEXPORTFORMAT']._serialized_start=13929 + _globals['_ANNOTATIONEXPORTFORMAT']._serialized_end=14038 _globals['_GETLATESTLOGGERDATAREQUEST']._serialized_start=46 _globals['_GETLATESTLOGGERDATAREQUEST']._serialized_end=303 _globals['_SIGNALCURVEINDEX']._serialized_start=306 @@ -141,122 +141,122 @@ _globals['_DATASTAT']._serialized_end=8018 _globals['_DATARECORD']._serialized_start=8020 _globals['_DATARECORD']._serialized_end=8082 - _globals['_DATASAMPLESRESPONSE']._serialized_start=8084 - _globals['_DATASAMPLESRESPONSE']._serialized_end=8174 - _globals['_HISTOGRAMSUBBAR']._serialized_start=8176 - _globals['_HISTOGRAMSUBBAR']._serialized_end=8243 - _globals['_HISTOGRAMBIN']._serialized_start=8245 - _globals['_HISTOGRAMBIN']._serialized_end=8349 - _globals['_CATEGORICALHISTOGRAMBAR']._serialized_start=8351 - _globals['_CATEGORICALHISTOGRAMBAR']._serialized_end=8442 - _globals['_HISTOGRAMREQUEST']._serialized_start=8444 - _globals['_HISTOGRAMREQUEST']._serialized_end=8496 - _globals['_HISTOGRAMRESPONSE']._serialized_start=8499 - _globals['_HISTOGRAMRESPONSE']._serialized_end=8677 - _globals['_GETMETADATAREQUEST']._serialized_start=8679 - _globals['_GETMETADATAREQUEST']._serialized_end=8766 - _globals['_GETMETADATARESPONSE']._serialized_start=8769 - _globals['_GETMETADATARESPONSE']._serialized_end=8922 - _globals['_STEPSAMPLESREQUEST']._serialized_start=8924 - _globals['_STEPSAMPLESREQUEST']._serialized_end=9030 - _globals['_STEPSAMPLESRESPONSE']._serialized_start=9032 - _globals['_STEPSAMPLESRESPONSE']._serialized_end=9155 - _globals['_GETSIGNALTRAJECTORYREQUEST']._serialized_start=9157 - _globals['_GETSIGNALTRAJECTORYREQUEST']._serialized_end=9246 - _globals['_SIGNALTRAJECTORY']._serialized_start=9248 - _globals['_SIGNALTRAJECTORY']._serialized_end=9300 - _globals['_GETSIGNALTRAJECTORYRESPONSE']._serialized_start=9302 - _globals['_GETSIGNALTRAJECTORYRESPONSE']._serialized_end=9427 - _globals['_POINTCLOUDREQUEST']._serialized_start=9429 - _globals['_POINTCLOUDREQUEST']._serialized_end=9518 - _globals['_POINTCLOUDCHUNK']._serialized_start=9521 - _globals['_POINTCLOUDCHUNK']._serialized_end=9712 - _globals['_MEDIAREQUEST']._serialized_start=9714 - _globals['_MEDIAREQUEST']._serialized_end=9812 - _globals['_MEDIACHUNK']._serialized_start=9815 - _globals['_MEDIACHUNK']._serialized_end=10089 - _globals['_DATAEDITSREQUEST']._serialized_start=10092 - _globals['_DATAEDITSREQUEST']._serialized_end=10360 - _globals['_DATAEDITSRESPONSE']._serialized_start=10362 - _globals['_DATAEDITSRESPONSE']._serialized_end=10415 - _globals['_DATASPLITSRESPONSE']._serialized_start=10417 - _globals['_DATASPLITSRESPONSE']._serialized_end=10475 - _globals['_AGENTHEALTHRESPONSE']._serialized_start=10477 - _globals['_AGENTHEALTHRESPONSE']._serialized_end=10534 - _globals['_INITIALIZEAGENTREQUEST']._serialized_start=10536 - _globals['_INITIALIZEAGENTREQUEST']._serialized_end=10630 - _globals['_INITIALIZEAGENTRESPONSE']._serialized_start=10632 - _globals['_INITIALIZEAGENTRESPONSE']._serialized_end=10691 - _globals['_CHANGEAGENTMODELREQUEST']._serialized_start=10693 - _globals['_CHANGEAGENTMODELREQUEST']._serialized_end=10733 - _globals['_CHANGEAGENTMODELRESPONSE']._serialized_start=10735 - _globals['_CHANGEAGENTMODELRESPONSE']._serialized_end=10795 - _globals['_GETAGENTMODELSREQUEST']._serialized_start=10797 - _globals['_GETAGENTMODELSREQUEST']._serialized_end=10820 - _globals['_GETAGENTMODELSRESPONSE']._serialized_start=10822 - _globals['_GETAGENTMODELSRESPONSE']._serialized_end=10896 - _globals['_RESETAGENTRESPONSE']._serialized_start=10898 - _globals['_RESETAGENTRESPONSE']._serialized_end=10952 - _globals['_CLEARAGENTHISTORYRESPONSE']._serialized_start=10954 - _globals['_CLEARAGENTHISTORYRESPONSE']._serialized_end=11015 - _globals['_COMPACTAGENTHISTORYRESPONSE']._serialized_start=11017 - _globals['_COMPACTAGENTHISTORYRESPONSE']._serialized_end=11080 - _globals['_GETAGENTCONTEXTUSAGERESPONSE']._serialized_start=11083 - _globals['_GETAGENTCONTEXTUSAGERESPONSE']._serialized_end=11312 - _globals['_RESTORECHECKPOINTREQUEST']._serialized_start=11314 - _globals['_RESTORECHECKPOINTREQUEST']._serialized_end=11365 - _globals['_RESTORECHECKPOINTRESPONSE']._serialized_start=11367 - _globals['_RESTORECHECKPOINTRESPONSE']._serialized_end=11428 - _globals['_EXPERIMENTRUNINFO']._serialized_start=11431 - _globals['_EXPERIMENTRUNINFO']._serialized_end=11599 - _globals['_LISTEXPERIMENTRUNSREQUEST']._serialized_start=11601 - _globals['_LISTEXPERIMENTRUNSREQUEST']._serialized_end=11628 - _globals['_LISTEXPERIMENTRUNSRESPONSE']._serialized_start=11630 - _globals['_LISTEXPERIMENTRUNSRESPONSE']._serialized_end=11692 - _globals['_RENAMEEXPERIMENTRUNREQUEST']._serialized_start=11694 - _globals['_RENAMEEXPERIMENTRUNREQUEST']._serialized_end=11765 - _globals['_RENAMEEXPERIMENTRUNRESPONSE']._serialized_start=11767 - _globals['_RENAMEEXPERIMENTRUNRESPONSE']._serialized_end=11830 - _globals['_SETEXPERIMENTRUNNOTESREQUEST']._serialized_start=11832 - _globals['_SETEXPERIMENTRUNNOTESREQUEST']._serialized_end=11902 - _globals['_SETEXPERIMENTRUNNOTESRESPONSE']._serialized_start=11904 - _globals['_SETEXPERIMENTRUNNOTESRESPONSE']._serialized_end=11969 - _globals['_TRIGGEREVALUATIONREQUEST']._serialized_start=11971 - _globals['_TRIGGEREVALUATIONREQUEST']._serialized_end=12053 - _globals['_TRIGGEREVALUATIONRESPONSE']._serialized_start=12055 - _globals['_TRIGGEREVALUATIONRESPONSE']._serialized_end=12116 - _globals['_GETEVALUATIONSTATUSREQUEST']._serialized_start=12118 - _globals['_GETEVALUATIONSTATUSREQUEST']._serialized_end=12146 - _globals['_GETEVALUATIONSTATUSRESPONSE']._serialized_start=12149 - _globals['_GETEVALUATIONSTATUSRESPONSE']._serialized_end=12278 - _globals['_CANCELEVALUATIONREQUEST']._serialized_start=12280 - _globals['_CANCELEVALUATIONREQUEST']._serialized_end=12321 - _globals['_CANCELEVALUATIONRESPONSE']._serialized_start=12323 - _globals['_CANCELEVALUATIONRESPONSE']._serialized_end=12383 - _globals['_RUNNOTEBOOKCELLREQUEST']._serialized_start=12385 - _globals['_RUNNOTEBOOKCELLREQUEST']._serialized_end=12440 - _globals['_NOTEBOOKCELLDONE']._serialized_start=12442 - _globals['_NOTEBOOKCELLDONE']._serialized_end=12492 - _globals['_INTERRUPTNOTEBOOKCELLREQUEST']._serialized_start=12494 - _globals['_INTERRUPTNOTEBOOKCELLREQUEST']._serialized_end=12524 - _globals['_INTERRUPTNOTEBOOKCELLRESPONSE']._serialized_start=12526 - _globals['_INTERRUPTNOTEBOOKCELLRESPONSE']._serialized_end=12584 - _globals['_NOTEBOOKCELLCHUNK']._serialized_start=12587 - _globals['_NOTEBOOKCELLCHUNK']._serialized_end=12776 - _globals['_NOTEBOOKRESPONSE']._serialized_start=12778 - _globals['_NOTEBOOKRESPONSE']._serialized_end=12861 - _globals['_SAVENOTEBOOKREQUEST']._serialized_start=12863 - _globals['_SAVENOTEBOOKREQUEST']._serialized_end=12918 - _globals['_SAVENOTEBOOKRESPONSE']._serialized_start=12920 - _globals['_SAVENOTEBOOKRESPONSE']._serialized_end=12997 - _globals['_GENERATENOTEBOOKCODEREQUEST']._serialized_start=12999 - _globals['_GENERATENOTEBOOKCODEREQUEST']._serialized_end=13066 - _globals['_GENERATENOTEBOOKCODERESPONSE']._serialized_start=13068 - _globals['_GENERATENOTEBOOKCODERESPONSE']._serialized_end=13160 - _globals['_EXPORTANNOTATIONSREQUEST']._serialized_start=13162 - _globals['_EXPORTANNOTATIONSREQUEST']._serialized_end=13288 - _globals['_EXPORTANNOTATIONSRESPONSE']._serialized_start=13291 - _globals['_EXPORTANNOTATIONSRESPONSE']._serialized_end=13427 - _globals['_EXPERIMENTSERVICE']._serialized_start=13979 - _globals['_EXPERIMENTSERVICE']._serialized_end=16409 + _globals['_DATASAMPLESRESPONSE']._serialized_start=8085 + _globals['_DATASAMPLESRESPONSE']._serialized_end=8236 + _globals['_HISTOGRAMSUBBAR']._serialized_start=8238 + _globals['_HISTOGRAMSUBBAR']._serialized_end=8305 + _globals['_HISTOGRAMBIN']._serialized_start=8307 + _globals['_HISTOGRAMBIN']._serialized_end=8411 + _globals['_CATEGORICALHISTOGRAMBAR']._serialized_start=8413 + _globals['_CATEGORICALHISTOGRAMBAR']._serialized_end=8504 + _globals['_HISTOGRAMREQUEST']._serialized_start=8506 + _globals['_HISTOGRAMREQUEST']._serialized_end=8558 + _globals['_HISTOGRAMRESPONSE']._serialized_start=8561 + _globals['_HISTOGRAMRESPONSE']._serialized_end=8739 + _globals['_GETMETADATAREQUEST']._serialized_start=8741 + _globals['_GETMETADATAREQUEST']._serialized_end=8828 + _globals['_GETMETADATARESPONSE']._serialized_start=8831 + _globals['_GETMETADATARESPONSE']._serialized_end=8984 + _globals['_STEPSAMPLESREQUEST']._serialized_start=8986 + _globals['_STEPSAMPLESREQUEST']._serialized_end=9092 + _globals['_STEPSAMPLESRESPONSE']._serialized_start=9094 + _globals['_STEPSAMPLESRESPONSE']._serialized_end=9217 + _globals['_GETSIGNALTRAJECTORYREQUEST']._serialized_start=9219 + _globals['_GETSIGNALTRAJECTORYREQUEST']._serialized_end=9308 + _globals['_SIGNALTRAJECTORY']._serialized_start=9310 + _globals['_SIGNALTRAJECTORY']._serialized_end=9362 + _globals['_GETSIGNALTRAJECTORYRESPONSE']._serialized_start=9364 + _globals['_GETSIGNALTRAJECTORYRESPONSE']._serialized_end=9489 + _globals['_POINTCLOUDREQUEST']._serialized_start=9491 + _globals['_POINTCLOUDREQUEST']._serialized_end=9580 + _globals['_POINTCLOUDCHUNK']._serialized_start=9583 + _globals['_POINTCLOUDCHUNK']._serialized_end=9774 + _globals['_MEDIAREQUEST']._serialized_start=9776 + _globals['_MEDIAREQUEST']._serialized_end=9874 + _globals['_MEDIACHUNK']._serialized_start=9877 + _globals['_MEDIACHUNK']._serialized_end=10151 + _globals['_DATAEDITSREQUEST']._serialized_start=10154 + _globals['_DATAEDITSREQUEST']._serialized_end=10422 + _globals['_DATAEDITSRESPONSE']._serialized_start=10424 + _globals['_DATAEDITSRESPONSE']._serialized_end=10477 + _globals['_DATASPLITSRESPONSE']._serialized_start=10479 + _globals['_DATASPLITSRESPONSE']._serialized_end=10537 + _globals['_AGENTHEALTHRESPONSE']._serialized_start=10539 + _globals['_AGENTHEALTHRESPONSE']._serialized_end=10596 + _globals['_INITIALIZEAGENTREQUEST']._serialized_start=10598 + _globals['_INITIALIZEAGENTREQUEST']._serialized_end=10692 + _globals['_INITIALIZEAGENTRESPONSE']._serialized_start=10694 + _globals['_INITIALIZEAGENTRESPONSE']._serialized_end=10753 + _globals['_CHANGEAGENTMODELREQUEST']._serialized_start=10755 + _globals['_CHANGEAGENTMODELREQUEST']._serialized_end=10795 + _globals['_CHANGEAGENTMODELRESPONSE']._serialized_start=10797 + _globals['_CHANGEAGENTMODELRESPONSE']._serialized_end=10857 + _globals['_GETAGENTMODELSREQUEST']._serialized_start=10859 + _globals['_GETAGENTMODELSREQUEST']._serialized_end=10882 + _globals['_GETAGENTMODELSRESPONSE']._serialized_start=10884 + _globals['_GETAGENTMODELSRESPONSE']._serialized_end=10958 + _globals['_RESETAGENTRESPONSE']._serialized_start=10960 + _globals['_RESETAGENTRESPONSE']._serialized_end=11014 + _globals['_CLEARAGENTHISTORYRESPONSE']._serialized_start=11016 + _globals['_CLEARAGENTHISTORYRESPONSE']._serialized_end=11077 + _globals['_COMPACTAGENTHISTORYRESPONSE']._serialized_start=11079 + _globals['_COMPACTAGENTHISTORYRESPONSE']._serialized_end=11142 + _globals['_GETAGENTCONTEXTUSAGERESPONSE']._serialized_start=11145 + _globals['_GETAGENTCONTEXTUSAGERESPONSE']._serialized_end=11374 + _globals['_RESTORECHECKPOINTREQUEST']._serialized_start=11376 + _globals['_RESTORECHECKPOINTREQUEST']._serialized_end=11427 + _globals['_RESTORECHECKPOINTRESPONSE']._serialized_start=11429 + _globals['_RESTORECHECKPOINTRESPONSE']._serialized_end=11490 + _globals['_EXPERIMENTRUNINFO']._serialized_start=11493 + _globals['_EXPERIMENTRUNINFO']._serialized_end=11661 + _globals['_LISTEXPERIMENTRUNSREQUEST']._serialized_start=11663 + _globals['_LISTEXPERIMENTRUNSREQUEST']._serialized_end=11690 + _globals['_LISTEXPERIMENTRUNSRESPONSE']._serialized_start=11692 + _globals['_LISTEXPERIMENTRUNSRESPONSE']._serialized_end=11754 + _globals['_RENAMEEXPERIMENTRUNREQUEST']._serialized_start=11756 + _globals['_RENAMEEXPERIMENTRUNREQUEST']._serialized_end=11827 + _globals['_RENAMEEXPERIMENTRUNRESPONSE']._serialized_start=11829 + _globals['_RENAMEEXPERIMENTRUNRESPONSE']._serialized_end=11892 + _globals['_SETEXPERIMENTRUNNOTESREQUEST']._serialized_start=11894 + _globals['_SETEXPERIMENTRUNNOTESREQUEST']._serialized_end=11964 + _globals['_SETEXPERIMENTRUNNOTESRESPONSE']._serialized_start=11966 + _globals['_SETEXPERIMENTRUNNOTESRESPONSE']._serialized_end=12031 + _globals['_TRIGGEREVALUATIONREQUEST']._serialized_start=12033 + _globals['_TRIGGEREVALUATIONREQUEST']._serialized_end=12115 + _globals['_TRIGGEREVALUATIONRESPONSE']._serialized_start=12117 + _globals['_TRIGGEREVALUATIONRESPONSE']._serialized_end=12178 + _globals['_GETEVALUATIONSTATUSREQUEST']._serialized_start=12180 + _globals['_GETEVALUATIONSTATUSREQUEST']._serialized_end=12208 + _globals['_GETEVALUATIONSTATUSRESPONSE']._serialized_start=12211 + _globals['_GETEVALUATIONSTATUSRESPONSE']._serialized_end=12340 + _globals['_CANCELEVALUATIONREQUEST']._serialized_start=12342 + _globals['_CANCELEVALUATIONREQUEST']._serialized_end=12383 + _globals['_CANCELEVALUATIONRESPONSE']._serialized_start=12385 + _globals['_CANCELEVALUATIONRESPONSE']._serialized_end=12445 + _globals['_RUNNOTEBOOKCELLREQUEST']._serialized_start=12447 + _globals['_RUNNOTEBOOKCELLREQUEST']._serialized_end=12502 + _globals['_NOTEBOOKCELLDONE']._serialized_start=12504 + _globals['_NOTEBOOKCELLDONE']._serialized_end=12554 + _globals['_INTERRUPTNOTEBOOKCELLREQUEST']._serialized_start=12556 + _globals['_INTERRUPTNOTEBOOKCELLREQUEST']._serialized_end=12586 + _globals['_INTERRUPTNOTEBOOKCELLRESPONSE']._serialized_start=12588 + _globals['_INTERRUPTNOTEBOOKCELLRESPONSE']._serialized_end=12646 + _globals['_NOTEBOOKCELLCHUNK']._serialized_start=12649 + _globals['_NOTEBOOKCELLCHUNK']._serialized_end=12838 + _globals['_NOTEBOOKRESPONSE']._serialized_start=12840 + _globals['_NOTEBOOKRESPONSE']._serialized_end=12923 + _globals['_SAVENOTEBOOKREQUEST']._serialized_start=12925 + _globals['_SAVENOTEBOOKREQUEST']._serialized_end=12980 + _globals['_SAVENOTEBOOKRESPONSE']._serialized_start=12982 + _globals['_SAVENOTEBOOKRESPONSE']._serialized_end=13059 + _globals['_GENERATENOTEBOOKCODEREQUEST']._serialized_start=13061 + _globals['_GENERATENOTEBOOKCODEREQUEST']._serialized_end=13128 + _globals['_GENERATENOTEBOOKCODERESPONSE']._serialized_start=13130 + _globals['_GENERATENOTEBOOKCODERESPONSE']._serialized_end=13222 + _globals['_EXPORTANNOTATIONSREQUEST']._serialized_start=13224 + _globals['_EXPORTANNOTATIONSREQUEST']._serialized_end=13350 + _globals['_EXPORTANNOTATIONSRESPONSE']._serialized_start=13353 + _globals['_EXPORTANNOTATIONSRESPONSE']._serialized_end=13489 + _globals['_EXPERIMENTSERVICE']._serialized_start=14041 + _globals['_EXPERIMENTSERVICE']._serialized_end=16471 # @@protoc_insertion_point(module_scope) diff --git a/weightslab/reporting.py b/weightslab/reporting.py index 11983cab..3c0505f0 100644 --- a/weightslab/reporting.py +++ b/weightslab/reporting.py @@ -316,6 +316,14 @@ def compute_distribution_entries( if col is None: entries.append({"name": str(requested), "resolved": False}) continue + # A media column (e.g. media:pred_video on a video-generation run) holds + # descriptor JSON, not numbers -- pd.to_numeric would coerce it all to NaN + # and the card would wrongly claim "no numeric values". Flag it as media so + # the card can say what it actually is. (See _distribution_card_html.) + _ms = _media_store() + if _ms is not None and _ms.is_media_column(col): + entries.append({"name": str(requested), "column": col, "resolved": True, "is_media": True}) + continue try: values = pd.to_numeric(df[col], errors="coerce").dropna() except Exception: @@ -496,6 +504,85 @@ def _resolve_runs_map(checkpoint_manager) -> dict: return {} +def _media_store(): + """Lazy handle to weightslab.data.media_store (None if unavailable), so this + module needn't hard-depend on it and never fails to import when it's absent.""" + try: + from weightslab.data import media_store + return media_store + except Exception: + return None + + +def _poster_data_uri(poster: bytes) -> Optional[str]: + """Wrap poster bytes as a data: URI, sniffing PNG vs JPEG (posters are always + a still image regardless of the underlying media kind). None when empty.""" + if not poster: + return None + if poster[:8] == b"\x89PNG\r\n\x1a\n": + mime = "image/png" + elif poster[:2] == b"\xff\xd8": + mime = "image/jpeg" + elif poster[:6] in (b"GIF87a", b"GIF89a"): + mime = "image/gif" + else: + mime = "image/png" # sensible default; browsers sniff anyway + return f"data:{mime};base64,{base64.b64encode(poster).decode('ascii')}" + + +def compute_media_examples(df: Optional[pd.DataFrame], max_fields: int = 8, + max_examples: int = 6) -> list: + """Discover media columns (``media:``) in the sample dataframe and pull + a few poster frames per field from the in-process media_store, so a + video/image/audio-generation run's actual artifacts show up in the report + instead of being invisible. Returns ``[]`` for a run with no media (every + non-media use case is unchanged). Bounded by ``max_fields``/``max_examples`` + so it never scales with dataset size.""" + if df is None or getattr(df, "empty", True): + return [] + ms = _media_store() + if ms is None: + return [] + try: + media_cols = [c for c in df.columns if isinstance(c, str) and ms.is_media_column(c)] + except Exception: + return [] + examples: list = [] + for col in media_cols[:max_fields]: + field = ms.field_from_column(col) + try: + present = df[col].notna() + count = int(present.sum()) + except Exception: + continue + if count == 0: + continue + try: + ids = _sample_ids_for_mask(df, present, max_examples) + except Exception: + ids = [] + kind = "" + thumbnails = [] + for sid in ids: + try: + entry = ms.get(field, sid) + except Exception: + entry = None + if not entry: + continue + kind = kind or str(entry.get("kind") or "") + uri = _poster_data_uri(entry.get("poster") or b"") + if uri: + thumbnails.append({"sample_id": str(sid), "poster_uri": uri}) + examples.append({ + "field": field, + "kind": kind or "media", + "count": count, + "thumbnails": thumbnails, + }) + return examples + + def collect_report_context( root_log_dir, logger_q, @@ -575,6 +662,7 @@ def collect_report_context( "distributions": compute_distribution_entries(df, distributions, plt), "dataframe": compute_dataframe_stats(df), "loss_shape_tags": summarize_loss_shape_tags(df), + "media": compute_media_examples(df), "plotting_available": plt is not None, "runs": list(runs_map.values()), } @@ -726,6 +814,13 @@ def _distribution_card_html(entry: dict, block_id: str) -> str: f'
No column matching "{name}" ' 'was found in the dataset.
' ) + if entry.get("is_media"): + col = html.escape(str(entry.get("column") or name)) + return head + ( + f'
"{col}" is a media column ' + '(images/video/audio), not a numeric signal — see the Generated Media ' + 'section for its samples.
' + ) if not entry.get("n"): return head + ( '
No numeric values logged for this column yet.
' @@ -747,6 +842,54 @@ def _distribution_card_html(entry: dict, block_id: str) -> str: return body +def _media_section_html(media: list) -> str: + """The Generated Media section: poster thumbnails per media field (video/ + image/audio/...). Returns "" when there is no media, so non-media reports are + byte-for-byte unchanged. Uses inline styles with neutral (light/dark-safe) + colors so it needs no additions to the report's stylesheet.""" + if not media: + return "" + card_style = ("border:1px solid rgba(128,128,128,0.3);border-radius:10px;" + "padding:14px 16px;background:rgba(128,128,128,0.06);min-width:240px") + thumb_style = ("width:104px;height:104px;object-fit:cover;border-radius:8px;" + "background:rgba(128,128,128,0.15);border:1px solid rgba(128,128,128,0.25)") + cards = [] + for m in media: + field = html.escape(str(m.get("field") or "")) + kind = html.escape(str(m.get("kind") or "media")) + count = int(m.get("count") or 0) + thumbs = m.get("thumbnails") or [] + if thumbs: + thumbs_html = "".join( + f'
' + f'' + f'
' + f'#{html.escape(str(t["sample_id"]))}
' + for t in thumbs + ) + else: + thumbs_html = ('

Media attached, but no poster ' + 'frames are cached in this process to preview.

') + cards.append( + f'
' + f'
' + f'{field}' + f'{kind}' + f'{count:,} sample(s)' + f'
' + f'
{thumbs_html}
' + f'
' + ) + return ( + '
' + '

Generated Media

' + f'
{"".join(cards)}
' + '
' + ) + + def _distributions_section_html(distributions: list) -> str: """The optional Distributions section, e.g. "add a histogram of train_loss" (action_params={"distributions": ["train_loss"]}) -- omitted @@ -1633,6 +1776,8 @@ def _chartjs_script_tag() -> str: {signals_html} + {media_section_html} + {distributions_section_html}
@@ -1737,6 +1882,7 @@ def render_report(context: dict, output_path, narrative: Optional[str] = None) - root_log_dir=html.escape(context.get("root_log_dir", "")), narrative=narrative_html, signals_html=signals_html, + media_section_html=_media_section_html(context.get("media") or []), distributions_section_html=_distributions_section_html(context.get("distributions") or []), loss_shape_html=_loss_shape_section_html(context.get("loss_shape_tags") or []), dataframe_html=_dataframe_section_html(context.get("dataframe") or {}), diff --git a/weightslab/src.py b/weightslab/src.py index 928c149b..498b329c 100644 --- a/weightslab/src.py +++ b/weightslab/src.py @@ -5570,7 +5570,13 @@ def _ai_report_generation_result( _dm = get_dataframe() df = _dm.get_combined_df() if _dm is not None else None except Exception as _e: - logger.debug("ai_report_generation: no sample dataframe available (%s).", _e) + # Warn (not debug): if get_combined_df() raises — e.g. array/media proxy + # conversion choking on a video-generation dataframe — the report would + # otherwise silently degrade to "No sample dataframe data available yet" + # and read as broken for no visible reason. Surface it, keep going. + logger.warning( + "ai_report_generation: sample dataframe unavailable (%s); the report's " + "Dataset/Media sections will be empty.", _e, exc_info=True) df = None # The narrative comes from the live agent — the same LLM call the Studio diff --git a/weightslab/trainer/services/data_service.py b/weightslab/trainer/services/data_service.py index 0a436ec5..dd065d33 100755 --- a/weightslab/trainer/services/data_service.py +++ b/weightslab/trainer/services/data_service.py @@ -584,6 +584,10 @@ def __init__(self, ctx): ) self._is_filtered = False # Track if the current view is filtered/modified by user + # Last known FULL (unfiltered) row count, cached whenever GetDataSamples + # serves the unfiltered view, so the subview ribbon can report an accurate + # "X of Y" total even while a filter is active. See GetDataSamples. + self._last_full_count = 0 # logger.info("[DataService] Skipping expensive startup computations (aspect ratio, natural sort, signals).") # These should be triggered on-demand or run in background to avoid blocking training start. @@ -4990,7 +4994,31 @@ def GetDataSamples(self, request, context): try: # Process the request directly without deduplication logicj - return self._process_get_data_samples(request, context) + resp = self._process_get_data_samples(request, context) + # Stamp the server-authoritative subview state onto EVERY response so + # a fresh client (private window, no cached UI state) can render the + # "you are viewing a subview" warning ribbon + reset in both grid and + # list mode. Mirrors self._is_filtered, which is set on agent + # masks/filters/sorts and cleared on @reset. + try: + is_sub = bool(getattr(self, "_is_filtered", False)) + resp.is_subview = is_sub + view_df = getattr(self, "_all_datasets_df", None) + if view_df is not None: + resp.view_count = int(len(view_df)) + # Remember the full-dataset size whenever we serve the + # UNfiltered view, so we can still report an accurate total + # (for the "X of Y samples" ribbon) once a filter is applied + # and _all_datasets_df holds only the subview. Survives across + # a fresh client connecting to this same backend, since the + # backend serves the full view at least once at startup. + if not is_sub: + self._last_full_count = int(len(view_df)) + full = int(getattr(self, "_last_full_count", 0) or 0) + resp.total_count = full if full else int(len(view_df)) + except Exception: + logger.debug("GetDataSamples: could not stamp subview state", exc_info=True) + return resp except Exception as e: logger.error("Error in GetDataSamples: %s", str(e), exc_info=True) From 40ad072d1e525f58e8aceab6d35986a01b79ac64 Mon Sep 17 00:00:00 2001 From: GuillaumePELLUET Date: Wed, 26 Aug 2026 12:38:46 +0200 Subject: [PATCH 03/29] fix CI core dumped --- weightslab/opencode_binary.py | 26 ++++++++++++++++++++------ 1 file changed, 20 insertions(+), 6 deletions(-) diff --git a/weightslab/opencode_binary.py b/weightslab/opencode_binary.py index 35c2cc9c..9b9b4b42 100644 --- a/weightslab/opencode_binary.py +++ b/weightslab/opencode_binary.py @@ -25,6 +25,7 @@ from __future__ import annotations +import atexit import logging import os import platform @@ -338,11 +339,20 @@ def ensure_managed_binary(version: Optional[str] = None, def ensure_managed_binary_in_background(reason: str = "", logger: Optional[logging.Logger] = None) -> None: """Install OpenCode in a daemon thread if it isn't already present, logging - the install. Idempotent, best-effort, and non-blocking -- the caller (a CLI - launch or ``import weightslab``) never waits on the ~180 MB download. - - Respects WEIGHTSLAB_OPENCODE_AUTODOWNLOAD. Only the FIRST call per process - does anything; the rest return immediately. + the install. Idempotent, best-effort, and non-blocking -- a long-running + caller (a CLI launch or ``import weightslab`` in an app) never waits on the + ~180 MB download. + + A short-lived caller that exits right after import IS made to wait (see the + ``atexit`` hook below): a daemon thread still doing network/ssl I/O when + CPython starts tearing down interpreter state on exit is a known segfault + vector (use-after-free in the ssl/socket C extensions, not a catchable + Python exception) -- observed as `python -c "import weightslab"` dying with + "Segmentation fault (core dumped)" right after the import finished. Joining + at atexit -- which runs in the main thread before ``Py_Finalize`` begins + tearing down module/C-extension state -- closes that race. This only costs + time on the very first import on a machine; once the binary is cached, + ``find_managed_binary()`` above short-circuits and no thread is spawned. """ global _bg_started log = logger or _LOGGER @@ -369,4 +379,8 @@ def _run(): except Exception as exc: # pragma: no cover - best-effort log.debug("OpenCode background install failed: %s", exc) - threading.Thread(target=_run, name="opencode-install", daemon=True).start() + t = threading.Thread(target=_run, name="opencode-install", daemon=True) + t.start() + # Bounded by download_managed_binary()'s own per-candidate _DOWNLOAD_TIMEOUT, + # so this can't hang process exit indefinitely -- see the docstring above. + atexit.register(t.join) From 0ed0500b00ccf570a5d77ad9f89e481ee4c1bcf4 Mon Sep 17 00:00:00 2001 From: GuillaumePELLUET Date: Wed, 26 Aug 2026 13:10:38 +0200 Subject: [PATCH 04/29] fix stucked CI --- .github/workflows/ci.yml | 11 ++++++ .../trainer/services/notebook_service.py | 35 +++++++++++++++++-- 2 files changed, 43 insertions(+), 3 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 8526450b..3a90f58a 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -200,6 +200,11 @@ jobs: # Depends on the fast 3.11 install gate; `gate` is also a direct need so this # job can read run_ci (the matrix and main-only jobs run independently). needs: [ gate, install ] + # Safety net: without this, a hang falls back to GitHub's 360-minute + # default. pytest's own --timeout=600 below should already catch a stuck + # test, but this bounds the whole job even if some hang manages to dodge + # that (e.g. a shared resource -- see WEIGHTSLAB_NOTEBOOK_EXEC_TIMEOUT). + timeout-minutes: 30 steps: - name: Checkout repository uses: actions/checkout@v4 @@ -221,6 +226,12 @@ jobs: - name: Run unit tests - General run: | export WEIGHTSLAB_LOG_LEVEL="DEBUG" + # Bounds EmbeddedKernelBridge's per-cell wait (default: unbounded -- + # see notebook_service.py) so a dead/unresponsive embedded kernel in + # TestNotebookKernelEmbedded fails that one test fast instead of + # holding the kernel's lock forever and wedging every later test in + # the same contract-test class behind it. + export WEIGHTSLAB_NOTEBOOK_EXEC_TIMEOUT="60" # A per-test timeout guards against any regression that hangs a test. python -m pytest ./tests -v --timeout=600 diff --git a/weightslab/trainer/services/notebook_service.py b/weightslab/trainer/services/notebook_service.py index 3026ff69..de0d2e54 100644 --- a/weightslab/trainer/services/notebook_service.py +++ b/weightslab/trainer/services/notebook_service.py @@ -312,6 +312,24 @@ def configure_embedded_kernel(enabled: bool) -> None: _EMBED_ENABLED = bool(enabled) +def _embedded_exec_timeout() -> float | None: + """Seconds EmbeddedKernelBridge waits for a single cell to finish, or None + for unbounded (the product default -- see EmbeddedKernelBridge.__init__). + + Override via WEIGHTSLAB_NOTEBOOK_EXEC_TIMEOUT for environments (CI) where a + dead/unresponsive kernel failing fast matters more than letting a + legitimately long-running cell finish -- an unbounded wait there wedges + every later test against the same shared kernel behind the first hang. + """ + raw = os.environ.get("WEIGHTSLAB_NOTEBOOK_EXEC_TIMEOUT", "").strip() + if not raw: + return None + try: + return float(raw) + except ValueError: + return None + + def _ipykernel_available() -> bool: try: import ipykernel # noqa: F401 @@ -545,10 +563,20 @@ class EmbeddedKernelBridge: NotebookService's kernel-construction lock. """ - def __init__(self, connection_file: Path, startup_timeout: float = 30.0): + def __init__(self, connection_file: Path, startup_timeout: float = 30.0, + exec_timeout: float | None = None): from jupyter_client import BlockingKernelClient self._lock = threading.Lock() self._busy = False + # None (the default) preserves the product's actual contract: a cell + # may legitimately run for a long time (e.g. training) and is meant to + # be stopped via interrupt(), never by a hidden deadline. Only pass a + # finite value where a dead/unresponsive kernel failing fast matters + # more than that contract -- e.g. WEIGHTSLAB_NOTEBOOK_EXEC_TIMEOUT in + # CI (see _get_kernel()): an unbounded wait here would otherwise hold + # `self._lock` forever, wedging every later cell run against the same + # shared kernel behind it. + self._exec_timeout = exec_timeout self._client = BlockingKernelClient(connection_file=str(connection_file)) self._client.load_connection_file() self._client.start_channels() @@ -564,7 +592,7 @@ def _work(): self._busy = True try: reply = self._client.execute_interactive( - code, allow_stdin=False, timeout=None, + code, allow_stdin=False, timeout=self._exec_timeout, output_hook=lambda msg: _append_iopub_output( lambda kind, payload: q.put((kind, payload)), msg), ) @@ -894,7 +922,8 @@ def _get_kernel(self): connection_file = get_embedded_kernel_connection_file() if connection_file is not None: try: - self._kernel = EmbeddedKernelBridge(connection_file) + self._kernel = EmbeddedKernelBridge( + connection_file, exec_timeout=_embedded_exec_timeout()) logger.info( "NotebookService: attached to embedded Jupyter kernel (%s)", connection_file) From 9e9820a75ff70777848090606873a63a96cd7f21 Mon Sep 17 00:00:00 2001 From: GuillaumePELLUET Date: Wed, 26 Aug 2026 13:50:09 +0200 Subject: [PATCH 05/29] Deselect scale tests from the fast CI job; cap test job at 15 minutes test_logger_scale.py deliberately builds multi-million-row DuckDB fixtures to stress-test large-scale queries -- real, by-design heavy work that easily exceeds the per-test timeout on a shared runner. Its own marker registration already said "deselect with -m 'not scale'" but the CI job never actually did, so each run burned its 600s per-test timeout on these instead of skipping them, reading as a hang stuck at the same progress percentage across unrelated pushes. --- .github/workflows/ci.yml | 13 +++++++++++-- 1 file changed, 11 insertions(+), 2 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 3a90f58a..c6bafd1c 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -204,7 +204,7 @@ jobs: # default. pytest's own --timeout=600 below should already catch a stuck # test, but this bounds the whole job even if some hang manages to dodge # that (e.g. a shared resource -- see WEIGHTSLAB_NOTEBOOK_EXEC_TIMEOUT). - timeout-minutes: 30 + timeout-minutes: 15 steps: - name: Checkout repository uses: actions/checkout@v4 @@ -233,7 +233,16 @@ jobs: # the same contract-test class behind it. export WEIGHTSLAB_NOTEBOOK_EXEC_TIMEOUT="60" # A per-test timeout guards against any regression that hangs a test. - python -m pytest ./tests -v --timeout=600 + # -m "not scale": tests/backend/test_logger_scale.py deliberately + # builds a multi-million-row DuckDB fixture to stress-test large-scale + # queries (see its module docstring) -- real, by-design heavy work, + # not a hang, but easily 600s+ on a shared/throttled runner. Its own + # marker registration (pyproject.toml) already says "deselect with + # -m 'not scale'"; this job just never actually did. Each `scale` + # test hitting the per-test timeout instead of being deselected burns + # 10 minutes AND leaves the job stuck at the same progress % across + # unrelated pushes, which read as a hang. + python -m pytest ./tests -v --timeout=600 -m "not scale" # ── Agent smoke test on a pip-installed package ─────────────────────────── # Proves the Option-2 promise end-to-end: install weightslab into a CLEAN From 4728db246a22f54c6a5f7fa3c000b915f0f5ced6 Mon Sep 17 00:00:00 2001 From: GuillaumePELLUET Date: Wed, 26 Aug 2026 17:13:24 +0200 Subject: [PATCH 06/29] Update CHANGELOG - v2.0.1.dev0 --- CHANGELOG.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 845b7d87..c995e940 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1 +1 @@ -# Changelog - 2026-10-05 v1.5.1 (0) +# Changelog - 2026-08-26 v2.0.1.dev0 From 328e46ee31150d7b50490f6c1c6967bf95e6eb4c Mon Sep 17 00:00:00 2001 From: GuillaumePELLUET Date: Fri, 4 Sep 2026 17:14:11 +0200 Subject: [PATCH 07/29] =?UTF-8?q?refactor(examples):=20one=20import=20surf?= =?UTF-8?q?ace=20=E2=80=94=20`import=20weightslab=20as=20wl`?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Examples no longer reach into the package internals. Every `from weightslab.components.global_monitoring import (guard_training_context, guard_testing_context)` is gone; the call sites use `wl.guard_training_context` / `wl.guard_testing_context`, which are the very same singletons (the top-level names are lazy re-exports of that module, verified by identity). The Ultralytics trainers now have the same treatment: `WLAwareTrainer`, `WLAwareSegmentationTrainer`, `WLAwareDataset` and `WLAwareSegmentationDataset` are re-exported from the package root, so a YOLO script reads `trainer=wl.WLAwareTrainer` with no deep import. They go through the lazy export map rather than a plain import because `ultralytics` is an optional extra -- importing it eagerly would break `import weightslab` for everyone without it. `__getattr__` now turns that missing dependency into an actionable error ("needs the optional 'ultralytics' package: pip install 'weightslab[ultralytics]'") instead of a bare ModuleNotFoundError, and still raises AttributeError for names that genuinely do not exist. Notebooks are stripped of outputs and execution counts (23 notebooks verified clean). Edits are byte-surgical rather than an nbformat round-trip: these files carry different JSON indents (1 for Jupyter, 2 for Colab-saved) and CRLF, so a round-trip would have reformatted whole files and buried the real change. Every cell's source was diffed against HEAD to confirm nothing else moved. Docs and code samples follow the same path (`wl.guard_*`, `wl.WLAwareTrainer`), including the README wandb-migration diff and both AGENTS.md files so the next agent does not reintroduce the deep imports. Prose mentions of the bare names are left as they are -- those names remain importable from the package, so the text stays correct. Co-Authored-By: Claude Opus 5 (1M context) --- AGENTS.md | 8 +- README.md | 6 +- docs/examples/lightning/classification.rst | 4 +- docs/examples/pytorch/classification.rst | 4 +- docs/examples/pytorch/clustering.rst | 4 +- docs/examples/pytorch/detection.rst | 2 +- docs/examples/pytorch/generation.rst | 2 +- docs/examples/pytorch/segmentation.rst | 2 +- docs/examples/ultralytics/detection.rst | 4 +- docs/examples/usecases/lidar_detection.rst | 2 +- docs/pytorch_lightning.rst | 4 +- docs/segmentation_usecase.rst | 2 +- docs/ultralytics.rst | 6 +- docs/usecases.rst | 2 +- weightslab/AGENTS.md | 2 +- weightslab/__init__.py | 34 +- .../Lightning/wl-classification/main.py | 8 +- .../PyTorch/wl-ads-recommendation.ipynb | 8 +- .../Notebooks/PyTorch/wl-classification.ipynb | 446 +- .../Notebooks/PyTorch/wl-clustering.ipynb | 3639 +---------------- .../Notebooks/PyTorch/wl-detection.ipynb | 3475 +--------------- .../PyTorch/wl-fraud-detection.ipynb | 8 +- .../Notebooks/PyTorch/wl-segmentation.ipynb | 8 +- .../Notebooks/PyTorch/ws-classification.ipynb | 8 +- ...olo-on-brain-tumor-detection-dataset.ipynb | 1641 +------- ...olo-on-carparts-segmentation-dataset.ipynb | 1426 +------ ...n-construction-ppe-detection-dataset.ipynb | 1426 +------ ...s-yolo-on-crack-segmentation-dataset.ipynb | 1426 +------ ...ralytics-yolo-on-homeobjects-dataset.ipynb | 1426 +------ ...tics-yolo-on-kitti-detection-dataset.ipynb | 1426 +------ ...lytics-yolo-on-medical-pills-dataset.ipynb | 1426 +------ ...yolo-on-package-segmentation-dataset.ipynb | 1426 +------ ...mentation-loss-shapes-classification.ipynb | 7 +- ...mentation-loss-shapes-classification.ipynb | 7 +- .../PyTorch/wl-ads-recommendation/main.py | 8 +- .../verify_integration.py | 8 +- .../PyTorch/wl-classification/main.py | 8 +- .../examples/PyTorch/wl-clustering/main.py | 8 +- .../examples/PyTorch/wl-detection/main.py | 9 +- .../PyTorch/wl-fraud-detection/main.py | 8 +- .../wl-fraud-detection/verify_integration.py | 8 +- .../PyTorch/wl-image-generation/main.py | 8 +- .../examples/PyTorch/wl-segmentation/main.py | 9 +- .../PyTorch/wl-video-generation/main.py | 8 +- .../examples/Ultralytics/wl-detection/main.py | 3 +- .../Usecases/wl-2d-lidar-detection/main.py | 8 +- .../Usecases/wl-3d-lidar-detection/main.py | 8 +- .../main.py | 5 +- .../Usecases/wl-fashion-mnist-signals/main.py | 8 +- .../Usecases/ws-signals-mnist/main.py | 5 +- weightslab/integrations/ultralytics/README.md | 5 +- .../integrations/ultralytics/__init__.py | 3 +- 52 files changed, 227 insertions(+), 19225 deletions(-) diff --git a/AGENTS.md b/AGENTS.md index 97c449d7..bf549c8c 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -120,10 +120,10 @@ the global ledger (`weightslab/weightslab/backend/ledgers.py`, Conventions that matter for correctness: -- Wrap the train step in `with guard_training_context:` and eval in - `with guard_testing_context:` (from - `weightslab.components.global_monitoring`). This is how pause/resume and - train/test separation work — **skip it and pause/resume or stats will misbehave.** +- Wrap the train step in `with wl.guard_training_context:` and eval in + `with wl.guard_testing_context:` (both re-exported at package level; no deep + import). This is how pause/resume and train/test separation work — **skip it + and pause/resume or stats will misbehave.** - Use `model.get_age()` (steps actually trained; survives checkpoint reloads), not the raw loop counter. - `task_type` on the dataset/model selects rendering: `classification`, diff --git a/README.md b/README.md index 1d4c59cf..e02d58ee 100644 --- a/README.md +++ b/README.md @@ -244,8 +244,6 @@ the exact samples causing them — so you can fix your data, not just log it. -import wandb +import weightslab as wl -+from weightslab.components.global_monitoring import ( -+ guard_training_context, guard_testing_context) + +@wl.signal(name="byte_adjusted_loss", subscribe_to="loss/CE") +def byte_adjusted_loss(ctx): return ctx.subscribed_value / ctx.image_bytes @@ -285,7 +283,7 @@ the exact samples causing them — so you can fix your data, not just log it. for epoch in range(1, args.epochs + 1): model.train() for x, y in train_loader: -+ with guard_training_context: ++ with wl.guard_training_context: logits = model(x.to(device)) loss = criterion(logits, y.to(device)) optimizer.zero_grad(); loss.backward(); optimizer.step() @@ -298,7 +296,7 @@ the exact samples causing them — so you can fix your data, not just log it. model.eval() with torch.no_grad(): for x, y in test_loader: -+ with guard_testing_context: ++ with wl.guard_testing_context: accuracy.update(model(x.to(device)), y) - wandb.log({"test/acc": accuracy.compute().item(), "epoch": epoch}) + wl.save_signals(preds_raw=logits, targets=y, diff --git a/docs/examples/lightning/classification.rst b/docs/examples/lightning/classification.rst index acbde286..419ac837 100644 --- a/docs/examples/lightning/classification.rst +++ b/docs/examples/lightning/classification.rst @@ -55,7 +55,7 @@ so the module receives already-tracked objects: self.metric = metric def training_step(self, batch, batch_idx): - with guard_training_context: + with wl.guard_training_context: x, ids, y, _ = batch logits = self.model(x) preds = torch.argmax(logits, dim=1) @@ -64,7 +64,7 @@ so the module receives already-tracked objects: return loss.mean() def validation_step(self, batch, batch_idx): - with guard_testing_context: + with wl.guard_testing_context: x, ids, y, _ = batch logits = self.model(x) preds = torch.argmax(logits, dim=1) diff --git a/docs/examples/pytorch/classification.rst b/docs/examples/pytorch/classification.rst index 3053d262..2845bcb8 100644 --- a/docs/examples/pytorch/classification.rst +++ b/docs/examples/pytorch/classification.rst @@ -103,7 +103,7 @@ aggregate curve view. .. code-block:: python def train(loader, model, optimizer, criterion, device): - with guard_training_context: + with wl.guard_training_context: inputs, ids, targets, _ = next(loader) outputs = model(inputs.to(device)) loss_per_sample = criterion(outputs, targets.to(device), @@ -112,7 +112,7 @@ aggregate curve view. optimizer.step() def test(loader, model, criterion, metric, device): - with guard_testing_context, torch.no_grad(): + with wl.guard_testing_context, torch.no_grad(): for inputs, ids, targets, _ in loader: criterion(outputs, targets, batch_ids=ids) metric(preds, targets, batch_ids=ids) diff --git a/docs/examples/pytorch/clustering.rst b/docs/examples/pytorch/clustering.rst index 079cf28c..3e6e47c2 100644 --- a/docs/examples/pytorch/clustering.rst +++ b/docs/examples/pytorch/clustering.rst @@ -53,7 +53,7 @@ inspect the hardest negatives in the studio. .. code-block:: python def train(loader, model, optimizer, device): - with guard_training_context: + with wl.guard_training_context: images, uids, labels, _ = next(loader) embeddings = model(images.to(device)) triplet_loss = compute_triplet_loss(embeddings, labels) @@ -61,7 +61,7 @@ inspect the hardest negatives in the studio. optimizer.step() def evaluate(loader, model, device): - with guard_testing_context, torch.no_grad(): + with wl.guard_testing_context, torch.no_grad(): for images, uids, labels, _ in loader: embeddings = model(images.to(device)) ... diff --git a/docs/examples/pytorch/detection.rst b/docs/examples/pytorch/detection.rst index d0b6d71b..7f50ab44 100644 --- a/docs/examples/pytorch/detection.rst +++ b/docs/examples/pytorch/detection.rst @@ -90,7 +90,7 @@ IoU as a distribution overlaid on each image. .. code-block:: python - with guard_training_context: + with wl.guard_training_context: outputs = model(inputs) preds = decode_predictions(outputs.detach(), grid_size, conf_thresh) diff --git a/docs/examples/pytorch/generation.rst b/docs/examples/pytorch/generation.rst index 69c2845d..7c283b50 100644 --- a/docs/examples/pytorch/generation.rst +++ b/docs/examples/pytorch/generation.rst @@ -68,7 +68,7 @@ contains two groups of sample IDs. .. code-block:: python - with guard_training_context: + with wl.guard_training_context: [img1, img2], [uid1, uid2], [label1, label2], _ = next(train_loader) cls_out, recon_out = model([img1, img2]) diff --git a/docs/examples/pytorch/segmentation.rst b/docs/examples/pytorch/segmentation.rst index 68ec2b37..3035e486 100644 --- a/docs/examples/pytorch/segmentation.rst +++ b/docs/examples/pytorch/segmentation.rst @@ -70,7 +70,7 @@ the predicted mask. .. code-block:: python - with guard_training_context: + with wl.guard_training_context: outputs = model(inputs) combined = bce_sample(outputs, targets, batch_ids=ids) \ + dice_sample(outputs, targets, batch_ids=ids) diff --git a/docs/examples/ultralytics/detection.rst b/docs/examples/ultralytics/detection.rst index 14168e47..f22e28ee 100644 --- a/docs/examples/ultralytics/detection.rst +++ b/docs/examples/ultralytics/detection.rst @@ -21,7 +21,7 @@ the example lives. What the example does --------------------- -One ``YOLO.train(trainer=WLAwareTrainer, ...)`` call gives you: +One ``YOLO.train(trainer=wl.WLAwareTrainer, ...)`` call gives you: - Per-sample box / cls / dfl loss and live NMS overlay (train split). - Per-sample IoU and post-NMS overlay (val split). @@ -37,7 +37,7 @@ Integration in three lines wl.watch_or_edit(cfg, flag="hyperparameters", defaults=cfg, poll_interval=1.0) wl.serve() - YOLO("yolo11n.pt").train(trainer=WLAwareTrainer, data=..., workers=0, amp=False) + YOLO("yolo11n.pt").train(trainer=wl.WLAwareTrainer, data=..., workers=0, amp=False) wl.keep_serving() Required kwargs: ``workers=0`` (UID counter lives in the parent process), diff --git a/docs/examples/usecases/lidar_detection.rst b/docs/examples/usecases/lidar_detection.rst index b86b13a5..3cad8072 100644 --- a/docs/examples/usecases/lidar_detection.rst +++ b/docs/examples/usecases/lidar_detection.rst @@ -84,7 +84,7 @@ WeightsLab integration (identical to image detection) } # Training loop - with guard_training_context: + with wl.guard_training_context: points, ids, targets, _ = next(train_loader) outputs = model(points.to(device)) preds = decode_3d_predictions(outputs.detach()) diff --git a/docs/pytorch_lightning.rst b/docs/pytorch_lightning.rst index 5643983d..06eb42a2 100644 --- a/docs/pytorch_lightning.rst +++ b/docs/pytorch_lightning.rst @@ -36,7 +36,7 @@ LightningModule excerpt self.metric_wl = metric_wl def training_step(self, batch): - with guard_training_context: + with wl.guard_training_context: x, ids, y = batch logits = self.model(x) preds = torch.argmax(logits, dim=1) @@ -49,7 +49,7 @@ LightningModule excerpt return loss_batch.mean() def validation_step(self, batch): - with guard_testing_context: + with wl.guard_testing_context: x, ids, y = batch logits = self.model(x) preds = torch.argmax(logits, dim=1) diff --git a/docs/segmentation_usecase.rst b/docs/segmentation_usecase.rst index 5fe8b83b..af248878 100644 --- a/docs/segmentation_usecase.rst +++ b/docs/segmentation_usecase.rst @@ -133,7 +133,7 @@ lists so ordering lines up: dtype=torch.long, ) - with guard_training_context: + with wl.guard_training_context: inputs, ids, labels, _ = next(loader) outputs = model(inputs) # [B, C, H, W] batch_idx = _instance_batch_idx(labels) diff --git a/docs/ultralytics.rst b/docs/ultralytics.rst index 5b7c214a..3c66ac4d 100644 --- a/docs/ultralytics.rst +++ b/docs/ultralytics.rst @@ -32,13 +32,12 @@ Minimal integration import weightslab as wl from ultralytics import YOLO - from weightslab.integrations.ultralytics import WLAwareTrainer wl.watch_or_edit(cfg, flag="hyperparameters", defaults=cfg) wl.serve() YOLO("yolo11n.pt").train( - trainer=WLAwareTrainer, + trainer=wl.WLAwareTrainer, data="my_dataset.yaml", imgsz=640, epochs=100, @@ -152,7 +151,6 @@ End-to-end sequence import os, yaml, torch import weightslab as wl - from weightslab.integrations.ultralytics import WLAwareTrainer from ultralytics import YOLO # 1) Load config and register as live hyperparameters @@ -170,7 +168,7 @@ End-to-end sequence # 4) Train — WLAwareTrainer handles all WL wiring internally YOLO(cfg["model"]["name"]).train( - trainer=WLAwareTrainer, + trainer=wl.WLAwareTrainer, data=str(cfg["data_root"]), imgsz=cfg["image_size"], epochs=cfg.get("training_steps_to_do") or 1000, diff --git a/docs/usecases.rst b/docs/usecases.rst index 8de11fbb..93c6523c 100644 --- a/docs/usecases.rst +++ b/docs/usecases.rst @@ -98,7 +98,7 @@ Why ``reduction="none"``: .. code-block:: python - with guard_training_context: + with wl.guard_training_context: inputs, ids, labels = next(loader) outputs = model(inputs.to(device)) preds = outputs.argmax(dim=1, keepdim=True) diff --git a/weightslab/AGENTS.md b/weightslab/AGENTS.md index 283f31c6..f31ebd7b 100644 --- a/weightslab/AGENTS.md +++ b/weightslab/AGENTS.md @@ -288,7 +288,7 @@ The automatic `tag:loss_shape` tag (§3.6) uses these same primitives. | Paired/contrastive samples, group-level signals | `PyTorch/wl-generation` | `wl.save_group_signals`; dataset emits 2 rows per item via a `uids` metadata key. | | Reactive signals / custom loss-shape tagging | `Usecases/wl-classification-signals_shape_classification`, `Usecases/ws-signals-mnist` | §3.6; the latter is the minimal variant with no custom classifier. | | PyTorch Lightning | `Lightning/wl-classification` | Same `watch_or_edit` calls as plain PyTorch; guards wrap `training_step`/`validation_step` bodies; `Trainer(log_every_n_steps=0, enable_checkpointing=False, logger=False)`. | -| Ultralytics YOLO (detect/segment) | `Ultralytics/wl-detection` | Don't call `watch_or_edit` for model/optimizer/data/loss/metric — pass `trainer=WLAwareTrainer` (or `WLAwareSegmentationTrainer`) from `weightslab.integrations.ultralytics` to `YOLO(...).train(...)`. It wires everything via UL callbacks; you only watch the run config as `flag="hyperparameters"`. | +| Ultralytics YOLO (detect/segment) | `Ultralytics/wl-detection` | Don't call `watch_or_edit` for model/optimizer/data/loss/metric — pass `trainer=wl.WLAwareTrainer` (or `wl.WLAwareSegmentationTrainer`) to `YOLO(...).train(...)`. It wires everything via UL callbacks; you only watch the run config as `flag="hyperparameters"`. | ### 3.10 Verifying an integration headlessly diff --git a/weightslab/__init__.py b/weightslab/__init__.py index 759de5ea..17532a43 100644 --- a/weightslab/__init__.py +++ b/weightslab/__init__.py @@ -32,6 +32,23 @@ "seed_everything": (".utils.tools", "seed_everything"), "guard_training_context": (".components.global_monitoring", "guard_training_context"), "guard_testing_context": (".components.global_monitoring", "guard_testing_context"), + # Ultralytics integration, so a YOLO script needs no deep import: + # YOLO(...).train(trainer=wl.WLAwareTrainer, ...) + # Laziness is load-bearing here rather than just a speed-up: `ultralytics` is + # an optional extra, so importing this eagerly would break `import weightslab` + # for every user who does not have it installed. The module is resolved (and + # a missing extra reported, see _MISSING_EXTRA) only if a name is touched. + "WLAwareTrainer": (".integrations.ultralytics", "WLAwareTrainer"), + "WLAwareSegmentationTrainer": (".integrations.ultralytics", "WLAwareSegmentationTrainer"), + "WLAwareDataset": (".integrations.ultralytics", "WLAwareDataset"), + "WLAwareSegmentationDataset": (".integrations.ultralytics", "WLAwareSegmentationDataset"), +} + +# Lazy exports whose module needs a third-party package weightslab does not +# depend on. A bare "No module named 'ultralytics'" gives no hint that an extra +# exists for it, so name the install command instead. +_MISSING_EXTRA = { + ".integrations.ultralytics": ("ultralytics", "weightslab[ultralytics]"), } # Everything re-exported straight from .src (attribute name == export name). for _name in ( @@ -61,7 +78,17 @@ def __getattr__(name): # PEP 562 module-level lazy attribute access if target is None: raise AttributeError(f"module {__name__!r} has no attribute {name!r}") module_name, attr = target - module = importlib.import_module(module_name, __name__) + try: + module = importlib.import_module(module_name, __name__) + except ImportError as error: + extra = _MISSING_EXTRA.get(module_name) + if extra is None or extra[0] not in str(error): + raise + package, install = extra + raise ImportError( + f"weightslab.{name} needs the optional {package!r} package: " + f"pip install '{install}'" + ) from error value = getattr(module, attr) globals()[name] = value # cache so __getattr__ isn't hit again return value @@ -265,6 +292,11 @@ def _clean(v: str) -> str: "pointcloud_thumbnail", "pointcloud_boxes", + "WLAwareTrainer", + "WLAwareSegmentationTrainer", + "WLAwareDataset", + "WLAwareSegmentationDataset", + "_BANNER", "__version__", "__license__", diff --git a/weightslab/examples/Lightning/wl-classification/main.py b/weightslab/examples/Lightning/wl-classification/main.py index 206cc193..5363c6e5 100644 --- a/weightslab/examples/Lightning/wl-classification/main.py +++ b/weightslab/examples/Lightning/wl-classification/main.py @@ -14,10 +14,6 @@ from torchmetrics.classification import Accuracy from weightslab.examples.utils.baseline_models.pytorch.models import FashionCNN as CNN -from weightslab.components.global_monitoring import ( - guard_training_context, - guard_testing_context -) # ============================================================================= @@ -109,7 +105,7 @@ def forward(self, x): return self.model(x) def training_step(self, batch): - with guard_training_context: + with wl.guard_training_context: x, ids, y = batch logits = self(x) # forward pass preds = torch.argmax(logits, dim=1) @@ -127,7 +123,7 @@ def training_step(self, batch): return loss def validation_step(self, batch): - with guard_testing_context: + with wl.guard_testing_context: x, ids, y = batch logits = self(x) preds = torch.argmax(logits, dim=1) diff --git a/weightslab/examples/Notebooks/PyTorch/wl-ads-recommendation.ipynb b/weightslab/examples/Notebooks/PyTorch/wl-ads-recommendation.ipynb index 185cf654..70558094 100644 --- a/weightslab/examples/Notebooks/PyTorch/wl-ads-recommendation.ipynb +++ b/weightslab/examples/Notebooks/PyTorch/wl-ads-recommendation.ipynb @@ -100,10 +100,6 @@ "from tqdm.auto import tqdm\n", "\n", "import weightslab as wl\n", - "from weightslab.components.global_monitoring import (\n", - " guard_training_context,\n", - " guard_testing_context,\n", - ")\n", "\n", "logging.basicConfig(level=logging.ERROR)\n", "device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n", @@ -334,7 +330,7 @@ "outputs": [], "source": [ "def train(loader, model, optimizer, criterion, device):\n", - " with guard_training_context:\n", + " with wl.guard_training_context:\n", " inputs, ids, labels = next(loader)\n", " inputs, labels = inputs.to(device), labels.to(device)\n", " optimizer.zero_grad()\n", @@ -350,7 +346,7 @@ "def test(loader, model, criterion, metric, device, n_batches):\n", " losses = torch.tensor(0.0, device=device)\n", " for inputs, ids, labels in loader:\n", - " with guard_testing_context:\n", + " with wl.guard_testing_context:\n", " inputs, labels = inputs.to(device), labels.to(device)\n", " logits = model(inputs)\n", " preds = logits.argmax(dim=1, keepdim=True)\n", diff --git a/weightslab/examples/Notebooks/PyTorch/wl-classification.ipynb b/weightslab/examples/Notebooks/PyTorch/wl-classification.ipynb index 58a22e9d..bebbebbd 100644 --- a/weightslab/examples/Notebooks/PyTorch/wl-classification.ipynb +++ b/weightslab/examples/Notebooks/PyTorch/wl-classification.ipynb @@ -90,18 +90,7 @@ "id": "PVreJLYJIod6", "outputId": "9e77d75f-ef33-4807-a6c6-c69bb8e4a29a" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Looking in indexes: https://test.pypi.org/simple/, https://pypi.org/simple/\n", - "\u001b[31mERROR: Could not find a version that satisfies the requirement weightslab==1.3.3.dev6 (from versions: 0.0.0, 1.1.1, 1.1.2, 1.1.3.dev0, 1.1.3, 1.1.4, 1.1.5, 1.1.6, 1.1.7.dev3, 1.1.7, 1.1.8, 1.1.9, 1.2.0, 1.2.1.dev0, 1.2.1, 1.2.2, 1.2.3.dev0, 1.2.3, 1.2.4, 1.2.5, 1.2.6, 1.3.0.dev0, 1.3.0, 1.3.1.dev1, 1.3.1.dev2, 1.3.1.dev3, 1.3.1, 1.3.2.dev0, 1.3.2, 1.3.3.dev0, 1.3.3.dev1, 1.3.3.dev2, 1.3.3.dev3, 1.3.3.dev4, 1.3.3.dev5, 1.3.3, 1.9.1.dev1, 1.9.1.dev2, 1.9.1.dev3, 1.9.1.dev4, 20260427.dev1, 20260427.dev2, 20260605.dev1, 20260310110810, 20260310112657.dev2659810514, 20260310112936.dev2035376017, 20260310114507.dev36735734, 20260310122733.dev571756140, 20260310130048.dev500189762, 20260310142756.dev2389960633, 20260310172643.dev1790296826, 20260310184554.dev3572723963, 20260310190409.dev16171358, 20260311001628.dev82842326, 20260311001938.dev237981957, 20260311094631.dev2433054235, 20260311100556.dev2759443978, 20260311111421.dev1419930837, 20260311152850.dev375337088, 20260415152016.dev1529547259, 20260515094038.dev295368024, 20260515094810.dev295368024, 20260515100544.dev709462904, 20260518115533.dev1831822349, 20260518120723.dev3543420334, 20260518121158.dev2454375021, 20260518125331.dev2144318793, 20260518152946.dev2688573905, 20260518153709.dev457831340, 20260518154038.dev3380177004, 20260518154550.dev2597851767, 20260518160221.dev67736641, 20260518160351.dev500880749, 20260518160424.dev3971069761, 20260519121616.dev3380177004, 20260519122038.dev2371413592, 20260519153002.dev2989691760, 20260519160247.dev4080968838, 20260519160445.dev1656037640, 20260519160647.dev1656037640, 20260519161900.dev3430116169, 20260520150146.dev2702511536, 20260520151257.dev3739150100, 20260520151549.dev3739150100, 20260520160858.dev3843696446, 20260520161330.dev950128852, 20260521100142.dev3061163053, 20260521100237.dev3763458477, 20260521100304.dev2708183098, 20260521101422.dev1111427347, 20260521101942.dev1751890316, 20260521102210.dev697219115, 20260521112427.dev1815380353, 20260528101213.dev1156931848, 20260528101621.dev3002779928, 20260528102838.dev1225638335, 20260528104815.dev146316976, 20260528104912.dev2432105952, 20260528144750.dev228483024, 20260529124654.dev1088603485, 20260602102247.dev3028334291, 20260602133031.dev612843591, 20260603164444.dev1140145833, 20260604132920.dev4091502713, 20260604141626.dev180590797, 20260604154956.dev1777185457, 20260604161731.dev1492593462, 20260605101712.dev3201556731, 20260605102734.dev307817433, 20260605153744.dev1218078270, 20260605154211.dev1284003932, 20260605155425.dev2249789890, 20260608122954.dev616519929, 20260608130348.dev103073330, 20260608142704.dev1750894834, 20260608143725.dev279338996, 20260608144113.dev1776325482, 20260608151812.dev1493920450, 20260609145543.dev4094763614, 20260609150002.dev1824924853, 20260609150942.dev1916930876, 20260609151231.dev2007644555, 20260609164013.dev3320873074, 20260609170327.dev427086775, 20260609170903.dev4090827397, 20260609170910.dev878249345, 20260609171010.dev315170550, 20260610134434.dev2867115601, 20260610135809.dev2659845748, 20260610140822.dev3917850211, 20260610142006.dev1779011735, 20260610145358.dev827089089, 20260611100101.dev2435580864, 20260611122329.dev2942846670, 20260612094419.dev318051075, 20260612174112.dev1945936079, 20260612175512.dev558737864, 20260612175811.dev4293577391, 20260614170924.dev435351911, 20260617105222.dev3481250246, 20260617105337.dev4027879073, 20260617150202.dev2472499979, 20260617150338.dev3507587562, 20260618091144.dev3638457713, 20260618091310.dev2980634153, 20260618091437.dev666220125, 20260618161401.dev2672707490, 20260618161455.dev1062308657, 20260618161503.dev869421065, 20260619154839.dev3124849231, 20260619154914.dev1825884740, 20260619155132.dev3059317318, 20260619155148.dev2195359802, 20260619155831.dev2547505056, 20260619160048.dev2036037367, 20260619160133.dev887552290, 20260619164552.dev545225221, 20260620125803.dev57300911, 20260620155947.dev163709547, 20260620160335.dev978317758, 20260620161239.dev1614881684, 20260620163624.dev3186857944, 20260620163931.dev174614870, 20260620170427.dev937264444, 20260621140029.dev2748921241, 20260622131949.dev1993441264, 20260622132042.dev41563737, 20260626162054.dev815684860, 20260626162149.dev1517909447, 20260629095034.dev2864733358, 20260629095109.dev2365834914, 20260629100925.dev3295139208, 20260705162125.dev3118493333, 20260705162448.dev4054744437, 20260705162608.dev112194667, 20260705163739.dev337968460, 20260705163915.dev1702540858, 20260706102840.dev1428624302, 20260706103324.dev4063235919, 20260706103956.dev1761667072, 20260706105255.dev612849548, 20260706122808.dev397480003, 20260706123234.dev2929755776, 20260706124917.dev3042106597, 20260706132510.dev3579337194, 20260707110009.dev944740769, 20260707110041.dev949537649, 20260707130443.dev1726432854, 20260707130503.dev2999317701, 20260707130900.dev2393574598, 20260708143958.dev3881013472, 20260708144357.dev578891966, 20260708144457.dev658501659, 20260708144826.dev1757704307, 20260708144934.dev3945788796, 20260708145114.dev663145800, 20260710162317.dev3300605375, 20260710162608.dev3232863539)\u001b[0m\u001b[31m\n", - "\u001b[0m\u001b[31mERROR: No matching distribution found for weightslab==1.3.3.dev6\u001b[0m\u001b[31m\n", - "\u001b[0m" - ] - } - ], + "outputs": [], "source": [ "# Install WeightsLab. Colab already ships torch, torchvision, numpy, scikit-learn\n", "# and Pillow, so nothing extra is needed here.\n", @@ -129,37 +118,7 @@ "id": "-uAGOoqVFETj", "outputId": "1d806744-f7be-4da8-9a8c-4b68ae436a13" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "15/07/2026-14:06:00.794 INFO:root:setup_logging: WeightsLab logging initialized - Log file: /tmp/tmpn4he5yzg/weightslab_logs/weightslab_20260715_140600.log\n", - "15/07/2026-14:06:00.796 INFO:weightslab:: WeightsLab package initialized - Log level: INFO, Log to file: True\n", - "15/07/2026-14:06:00.797 INFO:weightslab:: \n", - "╭─────────────────────────────────────────────────────────────────────────────────────────────────────╮\n", - "│ │\n", - "│ \u001b[32m$$\\ $$\\ \u001b[0m $$\\ $$\\ $$\\ \u001b[31m$$\\ \u001b[0m $$\\ │\n", - "│ \u001b[32m$$ | $\\ $$ |\u001b[0m \\__| $$ | $$ | \u001b[31m$$ | \u001b[0m $$ | │\n", - "│ \u001b[32m$$ |$$$\\ $$ |\u001b[0m $$$$$$\\ $$\\ $$$$$$\\ $$$$$$$\\ $$$$$$\\ $$$$$$$\\ \u001b[31m$$ | \u001b[0m $$$$$$\\ $$$$$$$\\ │\n", - "│ \u001b[32m$$ $$ $$\\$$ |\u001b[0m$$ __$$\\ $$ |$$ __$$\\ $$ __$$\\ \\_$$ _| $$ _____|\u001b[31m$$ | \u001b[0m \\____$$\\ $$ __$$\\ │\n", - "│ \u001b[32m$$$$ _$$$$ |\u001b[0m$$$$$$$$ |$$ |$$ / $$ |$$ | $$ | $$ | \\$$$$$$\\ \u001b[31m$$ | \u001b[0m $$$$$$$ |$$ | $$ | │\n", - "│ \u001b[32m$$$ / \\$$$ |\u001b[0m$$ ____|$$ |$$ | $$ |$$ | $$ | $$ |$$\\ \\____$$\\ \u001b[31m$$ | \u001b[0m$$ __$$ |$$ | $$ | │\n", - "│ \u001b[32m$$ / \\$$ |\u001b[0m\\$$$$$$$\\ $$ |\\$$$$$$$ |$$ | $$ | \\$$$$ |$$$$$$$ |\u001b[31m$$$$$$$$\\ \u001b[0m\\$$$$$$$ |$$$$$$$ | │\n", - "│ \u001b[32m\\__/ \\__|\u001b[0m \\_______|\\__| \\____$$ |\\__| \\__| \\____/ \\_______/ \u001b[31m\\________|\u001b[0m \\_______|\\_______/ │\n", - "│ \u001b[32m \u001b[0m $$\\ $$ |\u001b[31m\u001b[0m │\n", - "│ \u001b[32m \u001b[0m \\$$$$$$ |\u001b[31m\u001b[0m │\n", - "│ \u001b[32m \u001b[0m \\______/\u001b[31m\u001b[0m │\n", - "│ │\n", - "│ Inspect - Edit - Evolve Neural Networks │\n", - "│ │\n", - "╰─── By GrayBx, v1.3.3.dev5 ──────────────────────────────────────────────────────────────────────────╯\n", - "\n", - "\n", - "Using device: cuda\n" - ] - } - ], + "outputs": [], "source": [ "import os\n", "import tempfile\n", @@ -175,10 +134,6 @@ "\n", "import weightslab as wl\n", "from weightslab.examples.utils.baseline_models.pytorch.models import FashionCNN as CNN\n", - "from weightslab.components.global_monitoring import (\n", - " guard_training_context,\n", - " guard_testing_context,\n", - ")\n", "\n", "logging.basicConfig(level=logging.ERROR)\n", "device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n", @@ -198,7 +153,7 @@ }, { "cell_type": "code", - "execution_count": 4, + "execution_count": null, "metadata": { "id": "NoYmyRYwFETk" }, @@ -242,7 +197,7 @@ }, { "cell_type": "code", - "execution_count": 11, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/" @@ -250,17 +205,7 @@ "id": "pHaJYoqNFETk", "outputId": "6bae5503-1a48-48c3-de01-8534009157a8" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "15/07/2026-14:23:06.251 INFO:weightslab.components.checkpoint_manager:__init__: CheckpointManager initialized at /tmp/weightslab_mnist_yonuysxx\n", - "15/07/2026-14:23:06.253 INFO:weightslab.src:watch_or_edit: Registered new checkpoint manager in ledger\n", - "Experiment logs -> /tmp/weightslab_mnist_yonuysxx\n" - ] - } - ], + "outputs": [], "source": [ "log_dir = tempfile.mkdtemp(prefix=\"weightslab_mnist_\")\n", "\n", @@ -334,7 +279,7 @@ }, { "cell_type": "code", - "execution_count": 6, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/" @@ -342,102 +287,7 @@ "id": "d7HMuDrCFETk", "outputId": "b5bbf0b2-9464-46ce-dc04-4841e15d0c97" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "15/07/2026-14:06:03.264 INFO:weightslab.backend.model_interface:__init__: Using checkpoint manager from ledger\n", - "15/07/2026-14:06:03.266 INFO:weightslab.components.checkpoint_manager:load_checkpoint: Loading checkpoint 124b3608e98f2d94...\n", - "15/07/2026-14:06:03.267 INFO:weightslab.components.checkpoint_manager:load_checkpoint: Target: HP=124b3608 MODEL=e98f2d94 DATA=9e70c5a5\n", - "15/07/2026-14:06:03.268 INFO:weightslab.components.checkpoint_manager:load_checkpoint: Current: HP=124b3608 MODEL=e98f2d94 DATA=9e70c5a5\n", - "15/07/2026-14:06:03.269 INFO:weightslab.components.checkpoint_manager:load_checkpoint: [-] Model architecture unchanged, using current model\n", - "15/07/2026-14:06:03.292 INFO:weightslab.components.checkpoint_manager:load_checkpoint: [OK] Loaded weights from step 6750 with RNG state\n", - "15/07/2026-14:06:03.294 INFO:weightslab.components.checkpoint_manager:load_checkpoint: [-] Config unchanged, using current config\n", - "15/07/2026-14:06:03.296 INFO:weightslab.components.checkpoint_manager:load_checkpoint: Loaded components: {'weights'}\n", - "15/07/2026-14:06:03.300 INFO:weightslab.backend.model_interface:__init__: Auto-loaded model weights from checkpoint 124b3608e98f2d94 (step 6750)\n", - "15/07/2026-14:06:03.304 INFO:weightslab.data.data_samples_with_ops:__init__: [DataSampleTrackingWrapper] H5 persistence enabled at /tmp/weightslab_mnist_kg8215pb/checkpoints/data/data.h5\n", - "15/07/2026-14:06:03.402 INFO:weightslab.data.data_samples_with_ops:__init__: Preloading sample statistics for PHYSICAL indices in split 'train_loader' with: preload_labels=True,\n", - "15/07/2026-14:06:03.404 INFO:weightslab.data.data_samples_with_ops:__init__: Metadata will be loaded on demand from the wrapped dataset for split 'train_loader' when accessed, which may increase latency on first access but reduces initialization time.\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "Initializing ledger for split 'train_loader': 100%|██████████| 60000/60000 [00:12<00:00, 4884.89it/s]" - ] - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "15/07/2026-14:06:15.708 INFO:weightslab.data.dataframe_manager:register_split: [LedgeredDataFrameManager] Registering split 'train_loader' with 60000 samples.\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "\n" - ] - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "15/07/2026-14:06:16.062 INFO:weightslab.data.dataframe_manager:register_split: [LedgeredDataFrameManager] After annotation expansion: 60000 samples → 60000 annotation rows.\n", - "15/07/2026-14:06:16.064 INFO:weightslab.data.h5_array_store:__init__: [H5ArrayStore] Initialized with cache limit: 2048MB\n", - "15/07/2026-14:06:17.387 INFO:weightslab.components.checkpoint_manager:load_checkpoint: Loading checkpoint 124b3608e98f2d94...\n", - "15/07/2026-14:06:17.389 INFO:weightslab.components.checkpoint_manager:load_checkpoint: Target: HP=124b3608 MODEL=e98f2d94 DATA=9e70c5a5\n", - "15/07/2026-14:06:17.390 INFO:weightslab.components.checkpoint_manager:load_checkpoint: Current: HP=124b3608 MODEL=e98f2d94 DATA=9e70c5a5\n", - "15/07/2026-14:06:17.393 INFO:weightslab.components.checkpoint_manager:load_checkpoint: [-] Model architecture unchanged, using current model\n", - "15/07/2026-14:06:17.393 INFO:weightslab.components.checkpoint_manager:load_checkpoint: [-] Config unchanged, using current config\n", - "15/07/2026-14:06:17.593 INFO:weightslab.components.checkpoint_manager:load_checkpoint: [OK] Loaded data snapshot (70000 rows) with RNG state\n", - "15/07/2026-14:06:17.594 INFO:weightslab.components.checkpoint_manager:load_checkpoint: Loaded components: {'data'}\n", - "15/07/2026-14:06:17.918 INFO:weightslab.backend.dataloader_interface:_load_checkpoint_data: Applied data snapshot from checkpoint (70000 rows)\n", - "15/07/2026-14:06:17.922 INFO:weightslab.data.data_samples_with_ops:__init__: [DataSampleTrackingWrapper] H5 persistence enabled at /tmp/weightslab_mnist_kg8215pb/checkpoints/data/data.h5\n", - "15/07/2026-14:06:17.932 INFO:weightslab.data.data_samples_with_ops:__init__: Preloading sample statistics for PHYSICAL indices in split 'test_loader' with: preload_labels=True,\n", - "15/07/2026-14:06:17.935 INFO:weightslab.data.data_samples_with_ops:__init__: Metadata will be loaded on demand from the wrapped dataset for split 'test_loader' when accessed, which may increase latency on first access but reduces initialization time.\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "Initializing ledger for split 'test_loader': 100%|██████████| 10000/10000 [00:01<00:00, 7865.73it/s]" - ] - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "15/07/2026-14:06:19.213 INFO:weightslab.data.dataframe_manager:register_split: [LedgeredDataFrameManager] Registering split 'test_loader' with 10000 samples.\n", - "15/07/2026-14:06:19.273 INFO:weightslab.data.dataframe_manager:register_split: [LedgeredDataFrameManager] After annotation expansion: 10000 samples → 10000 annotation rows.\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "\n" - ] - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "15/07/2026-14:06:19.798 INFO:weightslab.components.checkpoint_manager:load_checkpoint: Loading checkpoint 124b3608e98f2d94...\n", - "15/07/2026-14:06:19.799 INFO:weightslab.components.checkpoint_manager:load_checkpoint: Target: HP=124b3608 MODEL=e98f2d94 DATA=9e70c5a5\n", - "15/07/2026-14:06:19.802 INFO:weightslab.components.checkpoint_manager:load_checkpoint: Current: HP=124b3608 MODEL=e98f2d94 DATA=9e70c5a5\n", - "15/07/2026-14:06:19.803 INFO:weightslab.components.checkpoint_manager:load_checkpoint: [-] Model architecture unchanged, using current model\n", - "15/07/2026-14:06:19.804 INFO:weightslab.components.checkpoint_manager:load_checkpoint: [-] Config unchanged, using current config\n", - "15/07/2026-14:06:19.965 INFO:weightslab.components.checkpoint_manager:load_checkpoint: [OK] Loaded data snapshot (70000 rows) with RNG state\n", - "15/07/2026-14:06:19.966 INFO:weightslab.components.checkpoint_manager:load_checkpoint: Loaded components: {'data'}\n", - "15/07/2026-14:06:20.277 INFO:weightslab.backend.dataloader_interface:_load_checkpoint_data: Applied data snapshot from checkpoint (70000 rows)\n" - ] - } - ], + "outputs": [], "source": [ "data_root = os.path.join(log_dir, \"data\")\n", "os.makedirs(data_root, exist_ok=True)\n", @@ -486,14 +336,14 @@ }, { "cell_type": "code", - "execution_count": 7, + "execution_count": null, "metadata": { "id": "LFFmdEMTFETl" }, "outputs": [], "source": [ "def train(loader, model, optimizer, criterion, device):\n", - " with guard_training_context:\n", + " with wl.guard_training_context:\n", " inputs, ids, labels = next(loader)\n", " inputs, labels = inputs.to(device), labels.to(device)\n", "\n", @@ -513,7 +363,7 @@ " losses = torch.tensor(0.0, device=device)\n", " st = time.time()\n", " for inputs, ids, labels in loader:\n", - " with guard_testing_context:\n", + " with wl.guard_testing_context:\n", " inputs, labels = inputs.to(device), labels.to(device)\n", " logits = model(inputs)\n", " preds = logits.argmax(dim=1, keepdim=True)\n", @@ -547,7 +397,7 @@ }, { "cell_type": "code", - "execution_count": 8, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/", @@ -556,109 +406,14 @@ "id": "YnEB98kvRBJt", "outputId": "590e5572-23b7-4bc1-8592-b666ce7f71dc" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "15/07/2026-14:06:20.303 INFO:weightslab.trainer.trainer_services:_run_security_preflight: \n", - "\n", - "# #######################################\n", - "# #######################################\n", - "[gRPC] Security preflight checks:\n", - "\tTLS: DISABLED (unencrypted traffic)\n", - "\tAuth tokens: NONE configured\n", - "\t! WARNING: GRPC_TLS_ENABLED=0. Traffic will be unencrypted. Use only for development.\n", - "\t! WARNING: No GRPC_AUTH_TOKEN/GRPC_AUTH_TOKENS configured. Only transport-level trust (TLS/mTLS) will protect RPC access.\n", - "# #######################################\n", - "# #######################################\n", - "\n", - "15/07/2026-14:06:20.305 INFO:weightslab.trainer.trainer_services:grpc_serve: [gRPC] Watchdogs disabled via WEIGHTSLAB_DISABLE_WATCHDOGS.\n", - "15/07/2026-14:06:20.306 INFO:weightslab.trainer.trainer_services:serving_thread_callback: [gRPC] Thread callback started\n", - "15/07/2026-14:06:20.307 INFO:weightslab.trainer.trainer_services:serving_thread_callback: [gRPC] Creating ThreadPoolExecutor with 6 worker threads (n_workers_grpc=None, max_concurrent_rpcs=None)\n", - "15/07/2026-14:06:20.308 INFO:weightslab.trainer.trainer_services:grpc_serve: [gRPC] Server started with watchdogs disabled (host=0.0.0.0 port=50051 workers=None)\n", - "15/07/2026-14:06:20.315 INFO:weightslab.trainer.trainer_services:serving_thread_callback: [gRPC] Server object created\n", - "15/07/2026-14:06:20.525 INFO:weightslab.trainer.services.agent.agent:__init__: Initializing DataManipulationAgent\n", - "15/07/2026-14:06:20.623 WARNING:weightslab.trainer.services.agent.agent:_load_config: Error loading config from /usr/local/lib/python3.12/dist-packages: [Errno 21] Is a directory: '/usr/local/lib/python3.12/dist-packages'\n", - "15/07/2026-14:06:20.625 INFO:weightslab.trainer.services.agent.agent:_load_config: \n", - "\n", - "# #######################################\n", - "# #######################################\n", - "Agent initialized from configuration /content/agent_config.yaml: \n", - "\tFinal Agent Configuration: Preferred Provider=openrouter, \n", - "\tFallback to Local=True, \n", - "\tOpenRouter Model=~google/gemini-flash-latest with:\n", - "\t\tAPI Key=None\n", - "\t\tBase URL=https://openrouter.ai/api/v1, \n", - "\tOllama Model=llama3.2:3b\n", - "# #######################################\n", - "# #######################################\n", - "\n", - "15/07/2026-14:06:20.625 INFO:weightslab.trainer.services.agent.agent:_setup_providers: Setting up Ollama with model llama3.2:3b\n", - "15/07/2026-14:06:20.710 INFO:weightslab.trainer.services.agent.agent:_setup_providers: [Agent] Ollama enabled: llama3.2:3b\n", - "15/07/2026-14:06:20.711 INFO:weightslab.trainer.services.data_service:_build_preview_cache: [PreviewCache] Building 64×64 or less or less preview cache for 2000 samples …\n", - "15/07/2026-14:06:20.712 INFO:weightslab.trainer.services.data_service:__init__: DataService initialized.\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "\r[PreviewCache]: 0%| | 0/2000 [00:00=1.24 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev7) (2.0.2)\n", - "Requirement already satisfied: pandas<3,>=2.2.2 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev7) (2.2.2)\n", - "Requirement already satisfied: duckdb<2,>=1.1 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev7) (1.3.2)\n", - "Requirement already satisfied: PyYAML<7,>=6.0.3 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev7) (6.0.3)\n", - "Requirement already satisfied: dill<0.5,>=0.3.8 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev7) (0.3.8)\n", - "Requirement already satisfied: zstandard<1,>=0.22 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev7) (0.25.0)\n", - "Requirement already satisfied: h5py<4,>=3.10 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev7) (3.16.0)\n", - "Requirement already satisfied: xxhash<4.1,>=3.4 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev7) (3.7.1)\n", - "Requirement already satisfied: tables<4,>=3.9 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev7) (3.10.2)\n", - "Collecting torch<=2.9,>=2.1 (from weightslab==1.3.3.dev7)\n", - " Downloading torch-2.9.0-cp312-cp312-manylinux_2_28_x86_64.whl.metadata (30 kB)\n", - "Requirement already satisfied: torchvision<1,>=0.16 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev7) (0.26.0+cpu)\n", - "Collecting torchmetrics>=1.9 (from weightslab==1.3.3.dev7)\n", - " Downloading torchmetrics-1.9.0-py3-none-any.whl.metadata (23 kB)\n", - "Requirement already satisfied: grpcio<2,>=1.80 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev7) (1.81.1)\n", - "Requirement already satisfied: protobuf<8,>=5.28.1 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev7) (5.29.6)\n", - "Requirement already satisfied: pydantic<3,>=2.7 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev7) (2.13.4)\n", - "Requirement already satisfied: Pillow<12,>=10 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev7) (11.3.0)\n", - "Requirement already satisfied: graphviz<1,>=0.20 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev7) (0.21)\n", - "Collecting onnx<=1.20,>=1.15 (from weightslab==1.3.3.dev7)\n", - " Downloading onnx-1.20.0-cp312-abi3-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl.metadata (8.4 kB)\n", - "Requirement already satisfied: tqdm<5,>=4.66 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev7) (4.67.3)\n", - "Requirement already satisfied: python-dotenv<2,>=1 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev7) (1.2.2)\n", - "Requirement already satisfied: langchain-core<2,>=0.3 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev7) (1.4.8)\n", - "Collecting langchain-ollama<2,>=0.2 (from weightslab==1.3.3.dev7)\n", - " Downloading langchain_ollama-1.1.0-py3-none-any.whl.metadata (3.0 kB)\n", - "Collecting langchain-openai<2,>=0.2 (from weightslab==1.3.3.dev7)\n", - " Downloading langchain_openai-1.3.5-py3-none-any.whl.metadata (3.4 kB)\n", - "Requirement already satisfied: typing-extensions~=4.12 in /usr/local/lib/python3.12/dist-packages (from grpcio<2,>=1.80->weightslab==1.3.3.dev7) (4.15.0)\n", - "Requirement already satisfied: jsonpatch<2.0.0,>=1.33.0 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev7) (1.33)\n", - "Requirement already satisfied: langchain-protocol>=0.0.17 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev7) (0.0.18)\n", - "Requirement already satisfied: langsmith<1.0.0,>=0.3.45 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev7) (0.9.1)\n", - "Requirement already satisfied: packaging>=23.2.0 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev7) (26.2)\n", - "Requirement already satisfied: tenacity!=8.4.0,<10.0.0,>=8.1.0 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev7) (9.1.4)\n", - "Requirement already satisfied: uuid-utils<1.0,>=0.12.0 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev7) (0.16.2)\n", - "Collecting ollama<1.0.0,>=0.6.1 (from langchain-ollama<2,>=0.2->weightslab==1.3.3.dev7)\n", - " Downloading ollama-0.6.2-py3-none-any.whl.metadata (5.8 kB)\n", - "Collecting langchain-core<2,>=0.3 (from weightslab==1.3.3.dev7)\n", - " Downloading langchain_core-1.4.9-py3-none-any.whl.metadata (4.7 kB)\n", - "Collecting openai<3.0.0,>=2.45.0 (from langchain-openai<2,>=0.2->weightslab==1.3.3.dev7)\n", - " Downloading openai-2.45.0-py3-none-any.whl.metadata (34 kB)\n", - "Requirement already satisfied: tiktoken<1.0.0,>=0.7.0 in /usr/local/lib/python3.12/dist-packages (from langchain-openai<2,>=0.2->weightslab==1.3.3.dev7) (0.13.0)\n", - "Requirement already satisfied: ml_dtypes>=0.5.0 in /usr/local/lib/python3.12/dist-packages (from onnx<=1.20,>=1.15->weightslab==1.3.3.dev7) (0.5.4)\n", - "Requirement already satisfied: python-dateutil>=2.8.2 in /usr/local/lib/python3.12/dist-packages (from pandas<3,>=2.2.2->weightslab==1.3.3.dev7) (2.9.0.post0)\n", - "Requirement already satisfied: pytz>=2020.1 in /usr/local/lib/python3.12/dist-packages (from pandas<3,>=2.2.2->weightslab==1.3.3.dev7) (2025.2)\n", - "Requirement already satisfied: tzdata>=2022.7 in /usr/local/lib/python3.12/dist-packages (from pandas<3,>=2.2.2->weightslab==1.3.3.dev7) (2026.2)\n", - "Requirement already satisfied: annotated-types>=0.6.0 in /usr/local/lib/python3.12/dist-packages (from pydantic<3,>=2.7->weightslab==1.3.3.dev7) (0.7.0)\n", - "Requirement already satisfied: pydantic-core==2.46.4 in /usr/local/lib/python3.12/dist-packages (from pydantic<3,>=2.7->weightslab==1.3.3.dev7) (2.46.4)\n", - "Requirement already satisfied: typing-inspection>=0.4.2 in /usr/local/lib/python3.12/dist-packages (from pydantic<3,>=2.7->weightslab==1.3.3.dev7) (0.4.2)\n", - "Requirement already satisfied: numexpr>=2.6.2 in /usr/local/lib/python3.12/dist-packages (from tables<4,>=3.9->weightslab==1.3.3.dev7) (2.14.1)\n", - "Requirement already satisfied: py-cpuinfo in /usr/local/lib/python3.12/dist-packages (from tables<4,>=3.9->weightslab==1.3.3.dev7) (9.0.0)\n", - "Requirement already satisfied: blosc2>=2.3.0 in /usr/local/lib/python3.12/dist-packages (from tables<4,>=3.9->weightslab==1.3.3.dev7) (4.5.1)\n", - "Requirement already satisfied: filelock in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7) (3.29.4)\n", - "Requirement already satisfied: setuptools in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7) (75.2.0)\n", - "Requirement already satisfied: sympy>=1.13.3 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7) (1.14.0)\n", - "Requirement already satisfied: networkx>=2.5.1 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7) (3.6.1)\n", - "Requirement already satisfied: jinja2 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7) (3.1.6)\n", - "Requirement already satisfied: fsspec>=0.8.5 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7) (2025.3.0)\n", - "Collecting nvidia-cuda-nvrtc-cu12==12.8.93 (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7)\n", - " Downloading nvidia_cuda_nvrtc_cu12-12.8.93-py3-none-manylinux2010_x86_64.manylinux_2_12_x86_64.whl.metadata (1.7 kB)\n", - "Collecting nvidia-cuda-runtime-cu12==12.8.90 (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7)\n", - " Downloading nvidia_cuda_runtime_cu12-12.8.90-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl.metadata (1.7 kB)\n", - "Collecting nvidia-cuda-cupti-cu12==12.8.90 (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7)\n", - " Downloading nvidia_cuda_cupti_cu12-12.8.90-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl.metadata (1.7 kB)\n", - "Collecting nvidia-cudnn-cu12==9.10.2.21 (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7)\n", - " Downloading nvidia_cudnn_cu12-9.10.2.21-py3-none-manylinux_2_27_x86_64.whl.metadata (1.8 kB)\n", - "Collecting nvidia-cublas-cu12==12.8.4.1 (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7)\n", - " Downloading nvidia_cublas_cu12-12.8.4.1-py3-none-manylinux_2_27_x86_64.whl.metadata (1.7 kB)\n", - "Collecting nvidia-cufft-cu12==11.3.3.83 (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7)\n", - " Downloading nvidia_cufft_cu12-11.3.3.83-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl.metadata (1.7 kB)\n", - "Collecting nvidia-curand-cu12==10.3.9.90 (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7)\n", - " Downloading nvidia_curand_cu12-10.3.9.90-py3-none-manylinux_2_27_x86_64.whl.metadata (1.7 kB)\n", - "Collecting nvidia-cusolver-cu12==11.7.3.90 (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7)\n", - " Downloading nvidia_cusolver_cu12-11.7.3.90-py3-none-manylinux_2_27_x86_64.whl.metadata (1.8 kB)\n", - "Collecting nvidia-cusparse-cu12==12.5.8.93 (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7)\n", - " Downloading nvidia_cusparse_cu12-12.5.8.93-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl.metadata (1.8 kB)\n", - "Collecting nvidia-cusparselt-cu12==0.7.1 (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7)\n", - " Downloading nvidia_cusparselt_cu12-0.7.1-py3-none-manylinux2014_x86_64.whl.metadata (7.0 kB)\n", - "Collecting nvidia-nccl-cu12==2.27.5 (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7)\n", - " Downloading nvidia_nccl_cu12-2.27.5-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl.metadata (2.0 kB)\n", - "Collecting nvidia-nvshmem-cu12==3.3.20 (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7)\n", - " Downloading nvidia_nvshmem_cu12-3.3.20-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl.metadata (2.1 kB)\n", - "Collecting nvidia-nvtx-cu12==12.8.90 (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7)\n", - " Downloading nvidia_nvtx_cu12-12.8.90-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl.metadata (1.8 kB)\n", - "Collecting nvidia-nvjitlink-cu12==12.8.93 (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7)\n", - " Downloading nvidia_nvjitlink_cu12-12.8.93-py3-none-manylinux2010_x86_64.manylinux_2_12_x86_64.whl.metadata (1.7 kB)\n", - "Collecting nvidia-cufile-cu12==1.13.1.3 (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7)\n", - " Downloading nvidia_cufile_cu12-1.13.1.3-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl.metadata (1.7 kB)\n", - "Collecting triton==3.5.0 (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7)\n", - " Downloading triton-3.5.0-cp312-cp312-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl.metadata (1.7 kB)\n", - "Collecting lightning-utilities>=0.15.3 (from torchmetrics>=1.9->weightslab==1.3.3.dev7)\n", - " Downloading lightning_utilities-0.15.3-py3-none-any.whl.metadata (5.5 kB)\n", - "INFO: pip is looking at multiple versions of torchvision to determine which version is compatible with other requirements. This could take a while.\n", - "Collecting torchvision<1,>=0.16 (from weightslab==1.3.3.dev7)\n", - " Downloading torchvision-0.28.0-cp312-cp312-manylinux_2_28_x86_64.whl.metadata (5.6 kB)\n", - " Downloading torchvision-0.27.1-cp312-cp312-manylinux_2_28_x86_64.whl.metadata (5.5 kB)\n", - " Downloading torchvision-0.27.0-cp312-cp312-manylinux_2_28_x86_64.whl.metadata (5.5 kB)\n", - " Downloading torchvision-0.26.0-cp312-cp312-manylinux_2_28_x86_64.whl.metadata (5.5 kB)\n", - " Downloading torchvision-0.25.0-cp312-cp312-manylinux_2_28_x86_64.whl.metadata (5.4 kB)\n", - " Downloading torchvision-0.24.1-cp312-cp312-manylinux_2_28_x86_64.whl.metadata (5.9 kB)\n", - " Downloading torchvision-0.24.0-cp312-cp312-manylinux_2_28_x86_64.whl.metadata (5.9 kB)\n", - "Requirement already satisfied: ndindex in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev7) (1.10.1)\n", - "Requirement already satisfied: msgpack in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev7) (1.2.1)\n", - "Requirement already satisfied: requests in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev7) (2.32.4)\n", - "Requirement already satisfied: rich in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev7) (13.9.4)\n", - "Requirement already satisfied: textual in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev7) (6.2.1)\n", - "Requirement already satisfied: textual-plotext in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev7) (1.0.1)\n", - "Requirement already satisfied: threadpoolctl in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev7) (3.6.0)\n", - "Requirement already satisfied: jsonpointer>=1.9 in /usr/local/lib/python3.12/dist-packages (from jsonpatch<2.0.0,>=1.33.0->langchain-core<2,>=0.3->weightslab==1.3.3.dev7) (3.1.1)\n", - "Requirement already satisfied: anyio>=3.5.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev7) (4.14.0)\n", - "Requirement already satisfied: distro>=1.7.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev7) (1.9.0)\n", - "Requirement already satisfied: httpx<1,>=0.23.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev7) (0.28.1)\n", - "Requirement already satisfied: orjson>=3.9.14 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev7) (3.11.9)\n", - "Requirement already satisfied: requests-toolbelt>=1.0.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev7) (1.0.0)\n", - "Requirement already satisfied: sniffio>=1.1 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev7) (1.3.1)\n", - "Requirement already satisfied: websockets>=15.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev7) (15.0.1)\n", - "Requirement already satisfied: jiter<1,>=0.10.0 in /usr/local/lib/python3.12/dist-packages (from openai<3.0.0,>=2.45.0->langchain-openai<2,>=0.2->weightslab==1.3.3.dev7) (0.15.0)\n", - "Requirement already satisfied: six>=1.5 in /usr/local/lib/python3.12/dist-packages (from python-dateutil>=2.8.2->pandas<3,>=2.2.2->weightslab==1.3.3.dev7) (1.17.0)\n", - "Requirement already satisfied: mpmath<1.4,>=1.1.0 in /usr/local/lib/python3.12/dist-packages (from sympy>=1.13.3->torch<=2.9,>=2.1->weightslab==1.3.3.dev7) (1.3.0)\n", - "Requirement already satisfied: regex in /usr/local/lib/python3.12/dist-packages (from tiktoken<1.0.0,>=0.7.0->langchain-openai<2,>=0.2->weightslab==1.3.3.dev7) (2025.11.3)\n", - "Requirement already satisfied: MarkupSafe>=2.0 in /usr/local/lib/python3.12/dist-packages (from jinja2->torch<=2.9,>=2.1->weightslab==1.3.3.dev7) (3.0.3)\n", - "Requirement already satisfied: idna>=2.8 in /usr/local/lib/python3.12/dist-packages (from anyio>=3.5.0->langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev7) (3.18)\n", - "Requirement already satisfied: certifi in /usr/local/lib/python3.12/dist-packages (from httpx<1,>=0.23.0->langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev7) (2026.6.17)\n", - "Requirement already satisfied: httpcore==1.* in /usr/local/lib/python3.12/dist-packages (from httpx<1,>=0.23.0->langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev7) (1.0.9)\n", - "Requirement already satisfied: h11>=0.16 in /usr/local/lib/python3.12/dist-packages (from httpcore==1.*->httpx<1,>=0.23.0->langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev7) (0.16.0)\n", - "Requirement already satisfied: charset_normalizer<4,>=2 in /usr/local/lib/python3.12/dist-packages (from requests->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev7) (3.4.7)\n", - "Requirement already satisfied: urllib3<3,>=1.21.1 in /usr/local/lib/python3.12/dist-packages (from requests->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev7) (2.5.0)\n", - "Requirement already satisfied: markdown-it-py>=2.2.0 in /usr/local/lib/python3.12/dist-packages (from rich->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev7) (4.2.0)\n", - "Requirement already satisfied: pygments<3.0.0,>=2.13.0 in /usr/local/lib/python3.12/dist-packages (from rich->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev7) (2.20.0)\n", - "Requirement already satisfied: platformdirs<5,>=3.6.0 in /usr/local/lib/python3.12/dist-packages (from textual->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev7) (4.10.0)\n", - "Requirement already satisfied: plotext<6.0.0,>=5.2.8 in /usr/local/lib/python3.12/dist-packages (from textual-plotext->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev7) (5.3.2)\n", - "Requirement already satisfied: mdurl~=0.1 in /usr/local/lib/python3.12/dist-packages (from markdown-it-py>=2.2.0->rich->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev7) (0.1.2)\n", - "Requirement already satisfied: linkify-it-py<3,>=1 in /usr/local/lib/python3.12/dist-packages (from markdown-it-py[linkify,plugins]>=2.1.0->textual->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev7) (2.1.0)\n", - "Requirement already satisfied: mdit-py-plugins>=0.5.0 in /usr/local/lib/python3.12/dist-packages (from markdown-it-py[linkify,plugins]>=2.1.0->textual->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev7) (0.6.1)\n", - "Requirement already satisfied: uc-micro-py in /usr/local/lib/python3.12/dist-packages (from linkify-it-py<3,>=1->markdown-it-py[linkify,plugins]>=2.1.0->textual->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev7) (2.0.0)\n", - "Downloading https://test-files.pythonhosted.org/packages/fd/50/2281a781f96c471a8e8ec9f7e15cb6b63bcae66b362474f00f1f29e75428/weightslab-1.3.3.dev7-py3-none-any.whl (2.8 MB)\n", - "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m2.8/2.8 MB\u001b[0m \u001b[31m16.9 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n", - "\u001b[?25hDownloading langchain_ollama-1.1.0-py3-none-any.whl (31 kB)\n", - "Downloading langchain_openai-1.3.5-py3-none-any.whl (121 kB)\n", - "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m121.6/121.6 kB\u001b[0m \u001b[31m3.1 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n", - "\u001b[?25hDownloading langchain_core-1.4.9-py3-none-any.whl (558 kB)\n", - "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m558.3/558.3 kB\u001b[0m \u001b[31m11.6 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n", - "\u001b[?25hDownloading onnx-1.20.0-cp312-abi3-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl (18.1 MB)\n", - "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m18.1/18.1 MB\u001b[0m \u001b[31m67.8 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n", - "\u001b[?25hDownloading torch-2.9.0-cp312-cp312-manylinux_2_28_x86_64.whl (899.7 MB)\n", - "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m899.7/899.7 MB\u001b[0m \u001b[31m1.1 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n", - "\u001b[?25hDownloading nvidia_cublas_cu12-12.8.4.1-py3-none-manylinux_2_27_x86_64.whl (594.3 MB)\n", - "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m594.3/594.3 MB\u001b[0m \u001b[31m3.1 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n", - "\u001b[?25hDownloading nvidia_cuda_cupti_cu12-12.8.90-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl (10.2 MB)\n", - "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m10.2/10.2 MB\u001b[0m \u001b[31m96.9 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n", - "\u001b[?25hDownloading nvidia_cuda_nvrtc_cu12-12.8.93-py3-none-manylinux2010_x86_64.manylinux_2_12_x86_64.whl (88.0 MB)\n", - "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m88.0/88.0 MB\u001b[0m \u001b[31m9.4 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n", - "\u001b[?25hDownloading nvidia_cuda_runtime_cu12-12.8.90-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl (954 kB)\n", - "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m954.8/954.8 kB\u001b[0m \u001b[31m64.7 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n", - "\u001b[?25hDownloading nvidia_cudnn_cu12-9.10.2.21-py3-none-manylinux_2_27_x86_64.whl (706.8 MB)\n", - "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m706.8/706.8 MB\u001b[0m \u001b[31m2.1 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n", - "\u001b[?25hDownloading nvidia_cufft_cu12-11.3.3.83-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl (193.1 MB)\n", - "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m193.1/193.1 MB\u001b[0m \u001b[31m6.2 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n", - "\u001b[?25hDownloading nvidia_cufile_cu12-1.13.1.3-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl (1.2 MB)\n", - "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m1.2/1.2 MB\u001b[0m \u001b[31m76.5 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n", - "\u001b[?25hDownloading nvidia_curand_cu12-10.3.9.90-py3-none-manylinux_2_27_x86_64.whl (63.6 MB)\n", - "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m63.6/63.6 MB\u001b[0m \u001b[31m12.6 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n", - "\u001b[?25hDownloading nvidia_cusolver_cu12-11.7.3.90-py3-none-manylinux_2_27_x86_64.whl (267.5 MB)\n", - "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m267.5/267.5 MB\u001b[0m \u001b[31m5.7 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n", - "\u001b[?25hDownloading nvidia_cusparse_cu12-12.5.8.93-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl (288.2 MB)\n", - "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m288.2/288.2 MB\u001b[0m \u001b[31m4.8 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n", - "\u001b[?25hDownloading nvidia_cusparselt_cu12-0.7.1-py3-none-manylinux2014_x86_64.whl (287.2 MB)\n", - "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m287.2/287.2 MB\u001b[0m \u001b[31m4.8 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n", - "\u001b[?25hDownloading nvidia_nccl_cu12-2.27.5-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl (322.3 MB)\n", - "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m322.3/322.3 MB\u001b[0m \u001b[31m1.5 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n", - "\u001b[?25hDownloading nvidia_nvjitlink_cu12-12.8.93-py3-none-manylinux2010_x86_64.manylinux_2_12_x86_64.whl (39.3 MB)\n", - "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m39.3/39.3 MB\u001b[0m \u001b[31m20.7 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n", - "\u001b[?25hDownloading nvidia_nvshmem_cu12-3.3.20-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl (124.7 MB)\n", - "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m124.7/124.7 MB\u001b[0m \u001b[31m7.5 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n", - "\u001b[?25hDownloading nvidia_nvtx_cu12-12.8.90-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl (89 kB)\n", - "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m90.0/90.0 kB\u001b[0m \u001b[31m9.8 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n", - "\u001b[?25hDownloading triton-3.5.0-cp312-cp312-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl (170.5 MB)\n", - "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m170.5/170.5 MB\u001b[0m \u001b[31m6.4 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n", - "\u001b[?25hDownloading torchmetrics-1.9.0-py3-none-any.whl (983 kB)\n", - "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m983.4/983.4 kB\u001b[0m \u001b[31m58.9 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n", - "\u001b[?25hDownloading torchvision-0.24.0-cp312-cp312-manylinux_2_28_x86_64.whl (8.1 MB)\n", - "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m8.1/8.1 MB\u001b[0m \u001b[31m82.8 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n", - "\u001b[?25hDownloading lightning_utilities-0.15.3-py3-none-any.whl (31 kB)\n", - "Downloading ollama-0.6.2-py3-none-any.whl (15 kB)\n", - "Downloading openai-2.45.0-py3-none-any.whl (1.6 MB)\n", - "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m1.6/1.6 MB\u001b[0m \u001b[31m64.8 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n", - "\u001b[?25hInstalling collected packages: nvidia-cusparselt-cu12, triton, nvidia-nvtx-cu12, nvidia-nvshmem-cu12, nvidia-nvjitlink-cu12, nvidia-nccl-cu12, nvidia-curand-cu12, nvidia-cufile-cu12, nvidia-cuda-runtime-cu12, nvidia-cuda-nvrtc-cu12, nvidia-cuda-cupti-cu12, nvidia-cublas-cu12, lightning-utilities, onnx, nvidia-cusparse-cu12, nvidia-cufft-cu12, nvidia-cudnn-cu12, openai, ollama, nvidia-cusolver-cu12, torch, langchain-core, torchvision, torchmetrics, langchain-openai, langchain-ollama, weightslab\n", - " Attempting uninstall: nvidia-nccl-cu12\n", - " Found existing installation: nvidia-nccl-cu12 2.30.7\n", - " Uninstalling nvidia-nccl-cu12-2.30.7:\n", - " Successfully uninstalled nvidia-nccl-cu12-2.30.7\n", - " Attempting uninstall: openai\n", - " Found existing installation: openai 2.43.0\n", - " Uninstalling openai-2.43.0:\n", - " Successfully uninstalled openai-2.43.0\n", - " Attempting uninstall: torch\n", - " Found existing installation: torch 2.11.0+cpu\n", - " Uninstalling torch-2.11.0+cpu:\n", - " Successfully uninstalled torch-2.11.0+cpu\n", - " Attempting uninstall: langchain-core\n", - " Found existing installation: langchain-core 1.4.8\n", - " Uninstalling langchain-core-1.4.8:\n", - " Successfully uninstalled langchain-core-1.4.8\n", - " Attempting uninstall: torchvision\n", - " Found existing installation: torchvision 0.26.0+cpu\n", - " Uninstalling torchvision-0.26.0+cpu:\n", - " Successfully uninstalled torchvision-0.26.0+cpu\n", - "Successfully installed langchain-core-1.4.9 langchain-ollama-1.1.0 langchain-openai-1.3.5 lightning-utilities-0.15.3 nvidia-cublas-cu12-12.8.4.1 nvidia-cuda-cupti-cu12-12.8.90 nvidia-cuda-nvrtc-cu12-12.8.93 nvidia-cuda-runtime-cu12-12.8.90 nvidia-cudnn-cu12-9.10.2.21 nvidia-cufft-cu12-11.3.3.83 nvidia-cufile-cu12-1.13.1.3 nvidia-curand-cu12-10.3.9.90 nvidia-cusolver-cu12-11.7.3.90 nvidia-cusparse-cu12-12.5.8.93 nvidia-cusparselt-cu12-0.7.1 nvidia-nccl-cu12-2.27.5 nvidia-nvjitlink-cu12-12.8.93 nvidia-nvshmem-cu12-3.3.20 nvidia-nvtx-cu12-12.8.90 ollama-0.6.2 onnx-1.20.0 openai-2.45.0 torch-2.9.0 torchmetrics-1.9.0 torchvision-0.24.0 triton-3.5.0 weightslab-1.3.3.dev7\n" - ] - } - ], + "outputs": [], "source": [ "# Install WeightsLab. Colab already ships torch, torchvision, numpy, scikit-learn\n", "# and Pillow, so nothing extra is needed here.\n", @@ -309,7 +91,7 @@ }, { "cell_type": "code", - "execution_count": 2, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/" @@ -317,38 +99,7 @@ "id": "kYbxEdSMF395", "outputId": "37ab580f-ae16-46ac-e394-9866f76469c7" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "16/07/2026-08:47:42.131 INFO:root:setup_logging: WeightsLab logging initialized - Log file: /tmp/tmpmfs9yakr/weightslab_logs/weightslab_20260716_084742.log\n", - "16/07/2026-08:47:42.132 INFO:weightslab:: WeightsLab package initialized - Log level: INFO, Log to file: True\n", - "16/07/2026-08:47:42.134 INFO:weightslab:: \n", - "╭─────────────────────────────────────────────────────────────────────────────────────────────────────╮\n", - "│ │\n", - "│ \u001b[32m$$\\ $$\\ \u001b[0m $$\\ $$\\ $$\\ \u001b[31m$$\\ \u001b[0m $$\\ │\n", - "│ \u001b[32m$$ | $\\ $$ |\u001b[0m \\__| $$ | $$ | \u001b[31m$$ | \u001b[0m $$ | │\n", - "│ \u001b[32m$$ |$$$\\ $$ |\u001b[0m $$$$$$\\ $$\\ $$$$$$\\ $$$$$$$\\ $$$$$$\\ $$$$$$$\\ \u001b[31m$$ | \u001b[0m $$$$$$\\ $$$$$$$\\ │\n", - "│ \u001b[32m$$ $$ $$\\$$ |\u001b[0m$$ __$$\\ $$ |$$ __$$\\ $$ __$$\\ \\_$$ _| $$ _____|\u001b[31m$$ | \u001b[0m \\____$$\\ $$ __$$\\ │\n", - "│ \u001b[32m$$$$ _$$$$ |\u001b[0m$$$$$$$$ |$$ |$$ / $$ |$$ | $$ | $$ | \\$$$$$$\\ \u001b[31m$$ | \u001b[0m $$$$$$$ |$$ | $$ | │\n", - "│ \u001b[32m$$$ / \\$$$ |\u001b[0m$$ ____|$$ |$$ | $$ |$$ | $$ | $$ |$$\\ \\____$$\\ \u001b[31m$$ | \u001b[0m$$ __$$ |$$ | $$ | │\n", - "│ \u001b[32m$$ / \\$$ |\u001b[0m\\$$$$$$$\\ $$ |\\$$$$$$$ |$$ | $$ | \\$$$$ |$$$$$$$ |\u001b[31m$$$$$$$$\\ \u001b[0m\\$$$$$$$ |$$$$$$$ | │\n", - "│ \u001b[32m\\__/ \\__|\u001b[0m \\_______|\\__| \\____$$ |\\__| \\__| \\____/ \\_______/ \u001b[31m\\________|\u001b[0m \\_______|\\_______/ │\n", - "│ \u001b[32m \u001b[0m $$\\ $$ |\u001b[31m\u001b[0m │\n", - "│ \u001b[32m \u001b[0m \\$$$$$$ |\u001b[31m\u001b[0m │\n", - "│ \u001b[32m \u001b[0m \\______/\u001b[31m\u001b[0m │\n", - "│ │\n", - "│ Inspect - Edit - Evolve Neural Networks │\n", - "│ │\n", - "╰─── By GrayBx, v1.3.3.dev7 ──────────────────────────────────────────────────────────────────────────╯\n", - "\n", - "\n", - "16/07/2026-08:47:42.137 INFO:weightslab.utils.telemetry:ping_import: WeightsLab uses anonymous usage data (package version used, and OS name). Set WL_NO_TELEMETRY=1 to disable.\n", - "Using device: cpu\n" - ] - } - ], + "outputs": [], "source": [ "import logging\n", "import os\n", @@ -365,10 +116,6 @@ "from typing import Any, Dict, List, Optional, Tuple\n", "\n", "import weightslab as wl\n", - "from weightslab.components.global_monitoring import (\n", - " guard_training_context,\n", - " guard_testing_context,\n", - ")\n", "\n", "logging.basicConfig(level=logging.ERROR)\n", "device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n", @@ -391,7 +138,7 @@ }, { "cell_type": "code", - "execution_count": 3, + "execution_count": null, "metadata": { "id": "94QCaUYjF396" }, @@ -641,7 +388,7 @@ }, { "cell_type": "code", - "execution_count": 4, + "execution_count": null, "metadata": { "id": "P6rEQuj6F397" }, @@ -860,7 +607,7 @@ }, { "cell_type": "code", - "execution_count": 5, + "execution_count": null, "metadata": { "id": "0mFrPy3QF397" }, @@ -1208,7 +955,7 @@ }, { "cell_type": "code", - "execution_count": 6, + "execution_count": null, "metadata": { "id": "zH13fDb8F398" }, @@ -1418,7 +1165,7 @@ }, { "cell_type": "code", - "execution_count": 7, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/" @@ -1426,56 +1173,7 @@ "id": "7f3X3gveF398", "outputId": "ccbe3f77-aefc-4d7d-d4be-4eea01f86dc9" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "16/07/2026-08:47:42.439 INFO:weightslab.src:watch_or_edit: LoggerQueue for experiment history has been initialized and registered.\n", - "16/07/2026-08:47:42.442 INFO:weightslab.components.checkpoint_manager:__init__: CheckpointManager initialized at /tmp/tmpu_lry8gs\n", - "16/07/2026-08:47:42.443 INFO:weightslab.src:watch_or_edit: Registered new checkpoint manager in ledger\n", - "16/07/2026-08:47:42.444 INFO:root:set_log_directory: Log file moved from /tmp/tmpmfs9yakr/weightslab_logs/weightslab_20260716_084742.log to /tmp/tmpu_lry8gs/weightslab_20260716_084742.log\n", - "16/07/2026-08:47:42.445 INFO:root:set_log_directory: Log directory updated to: /tmp/tmpu_lry8gs\n", - "16/07/2026-08:47:42.446 INFO:root:set_log_directory: Log file: /tmp/tmpu_lry8gs/weightslab_20260716_084742.log\n" - ] - }, - { - "data": { - "text/plain": [ - "{'experiment_name': 'face_triplet_training',\n", - " 'device': 'auto',\n", - " 'eval_full_to_train_steps_ratio': ValueProxy(key='eval_full_to_train_steps_ratio', value=50),\n", - " 'experiment_dump_to_train_steps_ratio': ValueProxy(key='experiment_dump_to_train_steps_ratio', value=10),\n", - " 'is_training': ValueProxy(key='is_training', value=True),\n", - " 'log_every': ValueProxy(key='log_every', value=50),\n", - " 'max_steps': ValueProxy(key='max_steps', value=1500),\n", - " 'enable_h5_persistence': ValueProxy(key='enable_h5_persistence', value=True),\n", - " 'serving_grpc': ValueProxy(key='serving_grpc', value=True),\n", - " 'serving_cli': ValueProxy(key='serving_cli', value=True),\n", - " 'data': {'dataset_type': 'olivetti',\n", - " 'root_data_dir': '',\n", - " 'image_size': 64,\n", - " 'train_ratio': 0.8,\n", - " 'min_images_per_class': 2,\n", - " 'train_loader': {'batch_size': 32, 'shuffle': True, 'n_workers': 0},\n", - " 'test_loader': {'batch_size': 64, 'shuffle': False, 'n_workers': 0}},\n", - " 'model': {'backbone': 'resnet18',\n", - " 'pretrained': True,\n", - " 'freeze_backbone': True,\n", - " 'embedding_dim': 128,\n", - " 'head_hidden_dim': 256,\n", - " 'lr': 0.001,\n", - " 'weight_decay': 0.0001,\n", - " 'loss': 'triplet',\n", - " 'margin': 0.3},\n", - " 'root_log_dir': '/tmp/tmpu_lry8gs'}" - ] - }, - "execution_count": 7, - "metadata": {}, - "output_type": "execute_result" - } - ], + "outputs": [], "source": [ "# ============================================================\n", "# Face Recognition - Training Config (inlined from config.yaml)\n", @@ -1566,7 +1264,7 @@ }, { "cell_type": "code", - "execution_count": 8, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/" @@ -1574,78 +1272,7 @@ "id": "91s0NMtEF399", "outputId": "07e8d642-4be4-48bb-86cc-aebfd17171f2" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "downloading Olivetti faces from https://ndownloader.figshare.com/files/5976027 to /root/scikit_learn_data\n", - "16/07/2026-08:47:46.760 INFO:__main__:__init__: FaceDataset [olivetti / train] | samples=320 | classes=40\n", - "16/07/2026-08:47:46.867 INFO:__main__:__init__: FaceDataset [olivetti / test] | samples=80 | classes=40\n", - "\n", - "Dataset : olivetti | train=320 | test=80 | classes=40\n", - "16/07/2026-08:47:46.880 INFO:weightslab.data.data_samples_with_ops:__init__: [DataSampleTrackingWrapper] H5 persistence enabled at /tmp/tmpu_lry8gs/checkpoints/data/data.h5\n", - "16/07/2026-08:47:46.893 INFO:weightslab.data.data_samples_with_ops:__init__: Preloading sample statistics for PHYSICAL indices in split 'train_loader' with: preload_labels=True, preload_metadata=True...\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "Initializing ledger for split 'train_loader': 0%| | 0/320 [00:00 the runtime device detected above).\n", @@ -1737,7 +1364,7 @@ }, { "cell_type": "code", - "execution_count": 9, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/" @@ -1745,44 +1372,7 @@ "id": "2J48fm6sF399", "outputId": "2970e5b8-1f7c-4dc4-ffd1-f8c7f50c5444" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Downloading: \"https://download.pytorch.org/models/resnet18-f37072fd.pth\" to /root/.cache/torch/hub/checkpoints/resnet18-f37072fd.pth\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "100%|██████████| 44.7M/44.7M [00:00<00:00, 77.5MB/s]\n" - ] - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "16/07/2026-08:47:49.494 INFO:__main__:_build_backbone: Backbone frozen ? training head only.\n", - "16/07/2026-08:47:49.505 INFO:weightslab.backend.model_interface:__init__: Using checkpoint manager from ledger\n", - "16/07/2026-08:47:49.509 INFO:__main__:__init__: FaceEmbeddingModel | backbone=resnet18 pretrained=True frozen=True | emb_dim=128 | trainable_params=164,736\n", - " Backbone : resnet18 (pretrained=True, frozen=True)\n", - " Emb dim : 128\n", - " Head dim : 256\n", - " Trainable : 164,736 params\n", - " Device : cpu\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "/usr/local/lib/python3.12/dist-packages/torch/nn/_reduction.py:51: UserWarning: size_average and reduce args will be deprecated, please use reduction='none' instead.\n", - " warnings.warn(warning.format(ret))\n" - ] - } - ], + "outputs": [], "source": [ "# ---- Model ----\n", "model_cfg = config.get(\"model\", {})\n", @@ -1815,7 +1405,7 @@ }, { "cell_type": "code", - "execution_count": 10, + "execution_count": null, "metadata": { "id": "1Cu9mwxEF399" }, @@ -1897,7 +1487,7 @@ " while True:\n", " step += 1\n", "\n", - " with guard_training_context:\n", + " with wl.guard_training_context:\n", " # ---- Fetch next batch (cycle loader) ----\n", " try:\n", " images, batch_ids, labels, _metadata = next(data_iter)\n", @@ -1936,7 +1526,7 @@ " )\n", "\n", " if should_eval:\n", - " with guard_testing_context:\n", + " with wl.guard_testing_context:\n", " print(f\"\\n[eval@test] step {step}\")\n", " metrics = evaluate(model=model, loader=test_loader, name=\"test\")\n", " eval_history.append({\"step\": step, \"metrics\": metrics})\n", @@ -1982,3202 +1572,7 @@ "id": "N_88-lktF39-", "outputId": "e736988d-2a95-4041-e9e5-414ec2861619" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "16/07/2026-08:47:49.574 INFO:weightslab.trainer.trainer_services:_run_security_preflight: \n", - "\n", - "# #######################################\n", - "# #######################################\n", - "[gRPC] Security preflight checks:\n", - "\tTLS: DISABLED (unencrypted traffic)\n", - "\tAuth tokens: NONE configured\n", - "\t! WARNING: GRPC_TLS_ENABLED=0. Traffic will be unencrypted. Use only for development.\n", - "\t! WARNING: No GRPC_AUTH_TOKEN/GRPC_AUTH_TOKENS configured. Only transport-level trust (TLS/mTLS) will protect RPC access.\n", - "# #######################################\n", - "# #######################################\n", - "\n", - "16/07/2026-08:47:49.576 INFO:weightslab.trainer.trainer_services:grpc_serve: [gRPC] Watchdogs disabled via WEIGHTSLAB_DISABLE_WATCHDOGS.\n", - "16/07/2026-08:47:49.577 INFO:weightslab.trainer.trainer_services:serving_thread_callback: [gRPC] Thread callback started\n", - "16/07/2026-08:47:49.580 INFO:weightslab.trainer.trainer_services:serving_thread_callback: [gRPC] Creating ThreadPoolExecutor with 6 worker threads (n_workers_grpc=None, max_concurrent_rpcs=None)\n", - "16/07/2026-08:47:49.578 INFO:weightslab.trainer.trainer_services:grpc_serve: [gRPC] Server started with watchdogs disabled (host=0.0.0.0 port=50051 workers=None)\n", - "16/07/2026-08:47:49.585 INFO:weightslab.tunnel:_ensure_bore: Downloading bore v0.6.0 (bore-v0.6.0-x86_64-unknown-linux-musl.tar.gz)...\n", - "16/07/2026-08:47:49.729 INFO:weightslab.trainer.trainer_services:serving_thread_callback: [gRPC] Server object created\n", - "16/07/2026-08:47:49.788 INFO:weightslab.trainer.services.agent.agent:__init__: Initializing DataManipulationAgent\n", - "16/07/2026-08:47:49.841 WARNING:weightslab.trainer.services.agent.agent:_load_config: Error loading config from /usr/local/lib/python3.12/dist-packages: [Errno 21] Is a directory: '/usr/local/lib/python3.12/dist-packages'\n", - "16/07/2026-08:47:49.853 INFO:weightslab.trainer.services.agent.agent:_load_config: \n", - "\n", - "# #######################################\n", - "# #######################################\n", - "Agent initialized from configuration /content/agent_config.yaml: \n", - "\tFinal Agent Configuration: Preferred Provider=openrouter, \n", - "\tFallback to Local=True, \n", - "\tOpenRouter Model=~google/gemini-flash-latest with:\n", - "\t\tAPI Key=None\n", - "\t\tBase URL=https://openrouter.ai/api/v1, \n", - "\tOllama Model=llama3.2:3b\n", - "# #######################################\n", - "# #######################################\n", - "\n", - "16/07/2026-08:47:49.860 INFO:weightslab.trainer.services.agent.agent:_setup_providers: Setting up Ollama with model llama3.2:3b\n", - "16/07/2026-08:47:50.171 INFO:weightslab.trainer.services.agent.agent:_setup_providers: [Agent] Ollama enabled: llama3.2:3b\n", - "16/07/2026-08:47:50.173 INFO:weightslab.trainer.services.data_service:_build_preview_cache: [PreviewCache] Building 64×64 or less or less preview cache for 400 samples …\n", - "16/07/2026-08:47:50.173 INFO:weightslab.trainer.services.data_service:__init__: DataService initialized.\n", - "16/07/2026-08:47:50.175 INFO:weightslab.trainer.trainer_services:serving_thread_callback: [gRPC] Servicer added\n", - "16/07/2026-08:47:50.176 INFO:weightslab.trainer.trainer_services:serving_thread_callback: [gRPC] Attempting to bind to 0.0.0.0:50051\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "\r[PreviewCache]: 0%| | 0/400 [00:00=1.24 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev7) (2.0.2)\n", - "Requirement already satisfied: pandas<3,>=2.2.2 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev7) (2.2.2)\n", - "Requirement already satisfied: duckdb<2,>=1.1 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev7) (1.3.2)\n", - "Requirement already satisfied: PyYAML<7,>=6.0.3 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev7) (6.0.3)\n", - "Requirement already satisfied: dill<0.5,>=0.3.8 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev7) (0.3.8)\n", - "Requirement already satisfied: zstandard<1,>=0.22 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev7) (0.25.0)\n", - "Requirement already satisfied: h5py<4,>=3.10 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev7) (3.16.0)\n", - "Requirement already satisfied: xxhash<4.1,>=3.4 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev7) (3.7.1)\n", - "Requirement already satisfied: tables<4,>=3.9 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev7) (3.10.2)\n", - "Requirement already satisfied: torch<=2.9,>=2.1 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev7) (2.9.0)\n", - "Requirement already satisfied: torchvision<1,>=0.16 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev7) (0.24.0)\n", - "Requirement already satisfied: torchmetrics>=1.9 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev7) (1.9.0)\n", - "Requirement already satisfied: grpcio<2,>=1.80 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev7) (1.81.1)\n", - "Requirement already satisfied: protobuf<8,>=5.28.1 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev7) (5.29.6)\n", - "Requirement already satisfied: pydantic<3,>=2.7 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev7) (2.13.4)\n", - "Requirement already satisfied: Pillow<12,>=10 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev7) (11.3.0)\n", - "Requirement already satisfied: graphviz<1,>=0.20 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev7) (0.21)\n", - "Requirement already satisfied: onnx<=1.20,>=1.15 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev7) (1.20.0)\n", - "Requirement already satisfied: tqdm<5,>=4.66 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev7) (4.67.3)\n", - "Requirement already satisfied: python-dotenv<2,>=1 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev7) (1.2.2)\n", - "Requirement already satisfied: langchain-core<2,>=0.3 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev7) (1.4.9)\n", - "Requirement already satisfied: langchain-ollama<2,>=0.2 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev7) (1.1.0)\n", - "Requirement already satisfied: langchain-openai<2,>=0.2 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev7) (1.3.5)\n", - "Requirement already satisfied: typing-extensions~=4.12 in /usr/local/lib/python3.12/dist-packages (from grpcio<2,>=1.80->weightslab==1.3.3.dev7) (4.15.0)\n", - "Requirement already satisfied: jsonpatch<2.0.0,>=1.33.0 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev7) (1.33)\n", - "Requirement already satisfied: langchain-protocol>=0.0.17 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev7) (0.0.18)\n", - "Requirement already satisfied: langsmith<1.0.0,>=0.3.45 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev7) (0.9.1)\n", - "Requirement already satisfied: packaging>=23.2.0 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev7) (26.2)\n", - "Requirement already satisfied: tenacity!=8.4.0,<10.0.0,>=8.1.0 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev7) (9.1.4)\n", - "Requirement already satisfied: uuid-utils<1.0,>=0.12.0 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev7) (0.16.2)\n", - "Requirement already satisfied: ollama<1.0.0,>=0.6.1 in /usr/local/lib/python3.12/dist-packages (from langchain-ollama<2,>=0.2->weightslab==1.3.3.dev7) (0.6.2)\n", - "Requirement already satisfied: openai<3.0.0,>=2.45.0 in /usr/local/lib/python3.12/dist-packages (from langchain-openai<2,>=0.2->weightslab==1.3.3.dev7) (2.45.0)\n", - "Requirement already satisfied: tiktoken<1.0.0,>=0.7.0 in /usr/local/lib/python3.12/dist-packages (from langchain-openai<2,>=0.2->weightslab==1.3.3.dev7) (0.13.0)\n", - "Requirement already satisfied: ml_dtypes>=0.5.0 in /usr/local/lib/python3.12/dist-packages (from onnx<=1.20,>=1.15->weightslab==1.3.3.dev7) (0.5.4)\n", - "Requirement already satisfied: python-dateutil>=2.8.2 in /usr/local/lib/python3.12/dist-packages (from pandas<3,>=2.2.2->weightslab==1.3.3.dev7) (2.9.0.post0)\n", - "Requirement already satisfied: pytz>=2020.1 in /usr/local/lib/python3.12/dist-packages (from pandas<3,>=2.2.2->weightslab==1.3.3.dev7) (2025.2)\n", - "Requirement already satisfied: tzdata>=2022.7 in /usr/local/lib/python3.12/dist-packages (from pandas<3,>=2.2.2->weightslab==1.3.3.dev7) (2026.2)\n", - "Requirement already satisfied: annotated-types>=0.6.0 in /usr/local/lib/python3.12/dist-packages (from pydantic<3,>=2.7->weightslab==1.3.3.dev7) (0.7.0)\n", - "Requirement already satisfied: pydantic-core==2.46.4 in /usr/local/lib/python3.12/dist-packages (from pydantic<3,>=2.7->weightslab==1.3.3.dev7) (2.46.4)\n", - "Requirement already satisfied: typing-inspection>=0.4.2 in /usr/local/lib/python3.12/dist-packages (from pydantic<3,>=2.7->weightslab==1.3.3.dev7) (0.4.2)\n", - "Requirement already satisfied: numexpr>=2.6.2 in /usr/local/lib/python3.12/dist-packages (from tables<4,>=3.9->weightslab==1.3.3.dev7) (2.14.1)\n", - "Requirement already satisfied: py-cpuinfo in /usr/local/lib/python3.12/dist-packages (from tables<4,>=3.9->weightslab==1.3.3.dev7) (9.0.0)\n", - "Requirement already satisfied: blosc2>=2.3.0 in /usr/local/lib/python3.12/dist-packages (from tables<4,>=3.9->weightslab==1.3.3.dev7) (4.5.1)\n", - "Requirement already satisfied: filelock in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7) (3.29.4)\n", - "Requirement already satisfied: setuptools in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7) (75.2.0)\n", - "Requirement already satisfied: sympy>=1.13.3 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7) (1.14.0)\n", - "Requirement already satisfied: networkx>=2.5.1 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7) (3.6.1)\n", - "Requirement already satisfied: jinja2 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7) (3.1.6)\n", - "Requirement already satisfied: fsspec>=0.8.5 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7) (2025.3.0)\n", - "Requirement already satisfied: nvidia-cuda-nvrtc-cu12==12.8.93 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7) (12.8.93)\n", - "Requirement already satisfied: nvidia-cuda-runtime-cu12==12.8.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7) (12.8.90)\n", - "Requirement already satisfied: nvidia-cuda-cupti-cu12==12.8.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7) (12.8.90)\n", - "Requirement already satisfied: nvidia-cudnn-cu12==9.10.2.21 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7) (9.10.2.21)\n", - "Requirement already satisfied: nvidia-cublas-cu12==12.8.4.1 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7) (12.8.4.1)\n", - "Requirement already satisfied: nvidia-cufft-cu12==11.3.3.83 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7) (11.3.3.83)\n", - "Requirement already satisfied: nvidia-curand-cu12==10.3.9.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7) (10.3.9.90)\n", - "Requirement already satisfied: nvidia-cusolver-cu12==11.7.3.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7) (11.7.3.90)\n", - "Requirement already satisfied: nvidia-cusparse-cu12==12.5.8.93 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7) (12.5.8.93)\n", - "Requirement already satisfied: nvidia-cusparselt-cu12==0.7.1 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7) (0.7.1)\n", - "Requirement already satisfied: nvidia-nccl-cu12==2.27.5 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7) (2.27.5)\n", - "Requirement already satisfied: nvidia-nvshmem-cu12==3.3.20 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7) (3.3.20)\n", - "Requirement already satisfied: nvidia-nvtx-cu12==12.8.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7) (12.8.90)\n", - "Requirement already satisfied: nvidia-nvjitlink-cu12==12.8.93 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7) (12.8.93)\n", - "Requirement already satisfied: nvidia-cufile-cu12==1.13.1.3 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7) (1.13.1.3)\n", - "Requirement already satisfied: triton==3.5.0 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev7) (3.5.0)\n", - "Requirement already satisfied: lightning-utilities>=0.15.3 in /usr/local/lib/python3.12/dist-packages (from torchmetrics>=1.9->weightslab==1.3.3.dev7) (0.15.3)\n", - "Requirement already satisfied: ndindex in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev7) (1.10.1)\n", - "Requirement already satisfied: msgpack in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev7) (1.2.1)\n", - "Requirement already satisfied: requests in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev7) (2.32.4)\n", - "Requirement already satisfied: rich in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev7) (13.9.4)\n", - "Requirement already satisfied: textual in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev7) (6.2.1)\n", - "Requirement already satisfied: textual-plotext in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev7) (1.0.1)\n", - "Requirement already satisfied: threadpoolctl in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev7) (3.6.0)\n", - "Requirement already satisfied: jsonpointer>=1.9 in /usr/local/lib/python3.12/dist-packages (from jsonpatch<2.0.0,>=1.33.0->langchain-core<2,>=0.3->weightslab==1.3.3.dev7) (3.1.1)\n", - "Requirement already satisfied: anyio>=3.5.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev7) (4.14.0)\n", - "Requirement already satisfied: distro>=1.7.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev7) (1.9.0)\n", - "Requirement already satisfied: httpx<1,>=0.23.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev7) (0.28.1)\n", - "Requirement already satisfied: orjson>=3.9.14 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev7) (3.11.9)\n", - "Requirement already satisfied: requests-toolbelt>=1.0.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev7) (1.0.0)\n", - "Requirement already satisfied: sniffio>=1.1 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev7) (1.3.1)\n", - "Requirement already satisfied: websockets>=15.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev7) (15.0.1)\n", - "Requirement already satisfied: jiter<1,>=0.10.0 in /usr/local/lib/python3.12/dist-packages (from openai<3.0.0,>=2.45.0->langchain-openai<2,>=0.2->weightslab==1.3.3.dev7) (0.15.0)\n", - "Requirement already satisfied: six>=1.5 in /usr/local/lib/python3.12/dist-packages (from python-dateutil>=2.8.2->pandas<3,>=2.2.2->weightslab==1.3.3.dev7) (1.17.0)\n", - "Requirement already satisfied: mpmath<1.4,>=1.1.0 in /usr/local/lib/python3.12/dist-packages (from sympy>=1.13.3->torch<=2.9,>=2.1->weightslab==1.3.3.dev7) (1.3.0)\n", - "Requirement already satisfied: regex in /usr/local/lib/python3.12/dist-packages (from tiktoken<1.0.0,>=0.7.0->langchain-openai<2,>=0.2->weightslab==1.3.3.dev7) (2025.11.3)\n", - "Requirement already satisfied: MarkupSafe>=2.0 in /usr/local/lib/python3.12/dist-packages (from jinja2->torch<=2.9,>=2.1->weightslab==1.3.3.dev7) (3.0.3)\n", - "Requirement already satisfied: idna>=2.8 in /usr/local/lib/python3.12/dist-packages (from anyio>=3.5.0->langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev7) (3.18)\n", - "Requirement already satisfied: certifi in /usr/local/lib/python3.12/dist-packages (from httpx<1,>=0.23.0->langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev7) (2026.6.17)\n", - "Requirement already satisfied: httpcore==1.* in /usr/local/lib/python3.12/dist-packages (from httpx<1,>=0.23.0->langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev7) (1.0.9)\n", - "Requirement already satisfied: h11>=0.16 in /usr/local/lib/python3.12/dist-packages (from httpcore==1.*->httpx<1,>=0.23.0->langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev7) (0.16.0)\n", - "Requirement already satisfied: charset_normalizer<4,>=2 in /usr/local/lib/python3.12/dist-packages (from requests->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev7) (3.4.7)\n", - "Requirement already satisfied: urllib3<3,>=1.21.1 in /usr/local/lib/python3.12/dist-packages (from requests->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev7) (2.5.0)\n", - "Requirement already satisfied: markdown-it-py>=2.2.0 in /usr/local/lib/python3.12/dist-packages (from rich->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev7) (4.2.0)\n", - "Requirement already satisfied: pygments<3.0.0,>=2.13.0 in /usr/local/lib/python3.12/dist-packages (from rich->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev7) (2.20.0)\n", - "Requirement already satisfied: platformdirs<5,>=3.6.0 in /usr/local/lib/python3.12/dist-packages (from textual->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev7) (4.10.0)\n", - "Requirement already satisfied: plotext<6.0.0,>=5.2.8 in /usr/local/lib/python3.12/dist-packages (from textual-plotext->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev7) (5.3.2)\n", - "Requirement already satisfied: mdurl~=0.1 in /usr/local/lib/python3.12/dist-packages (from markdown-it-py>=2.2.0->rich->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev7) (0.1.2)\n", - "Requirement already satisfied: linkify-it-py<3,>=1 in /usr/local/lib/python3.12/dist-packages (from markdown-it-py[linkify,plugins]>=2.1.0->textual->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev7) (2.1.0)\n", - "Requirement already satisfied: mdit-py-plugins>=0.5.0 in /usr/local/lib/python3.12/dist-packages (from markdown-it-py[linkify,plugins]>=2.1.0->textual->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev7) (0.6.1)\n", - "Requirement already satisfied: uc-micro-py in /usr/local/lib/python3.12/dist-packages (from linkify-it-py<3,>=1->markdown-it-py[linkify,plugins]>=2.1.0->textual->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev7) (2.0.0)\n" - ] - } - ], + "outputs": [], "source": [ "# Install WeightsLab. Colab already ships torch, torchvision, numpy, scikit-learn\n", "# and Pillow, so nothing extra is needed here.\n", @@ -203,7 +94,7 @@ }, { "cell_type": "code", - "execution_count": 2, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/" @@ -211,38 +102,7 @@ "id": "NHwcL67i_xBd", "outputId": "336670cc-0cab-4abf-b13f-91f6babe5c14" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "16/07/2026-08:32:40.383 INFO:root:setup_logging: WeightsLab logging initialized - Log file: /tmp/tmp_zk4n7mc/weightslab_logs/weightslab_20260716_083240.log\n", - "16/07/2026-08:32:40.387 INFO:weightslab:: WeightsLab package initialized - Log level: INFO, Log to file: True\n", - "16/07/2026-08:32:40.388 INFO:weightslab:: \n", - "╭─────────────────────────────────────────────────────────────────────────────────────────────────────╮\n", - "│ │\n", - "│ \u001b[32m$$\\ $$\\ \u001b[0m $$\\ $$\\ $$\\ \u001b[31m$$\\ \u001b[0m $$\\ │\n", - "│ \u001b[32m$$ | $\\ $$ |\u001b[0m \\__| $$ | $$ | \u001b[31m$$ | \u001b[0m $$ | │\n", - "│ \u001b[32m$$ |$$$\\ $$ |\u001b[0m $$$$$$\\ $$\\ $$$$$$\\ $$$$$$$\\ $$$$$$\\ $$$$$$$\\ \u001b[31m$$ | \u001b[0m $$$$$$\\ $$$$$$$\\ │\n", - "│ \u001b[32m$$ $$ $$\\$$ |\u001b[0m$$ __$$\\ $$ |$$ __$$\\ $$ __$$\\ \\_$$ _| $$ _____|\u001b[31m$$ | \u001b[0m \\____$$\\ $$ __$$\\ │\n", - "│ \u001b[32m$$$$ _$$$$ |\u001b[0m$$$$$$$$ |$$ |$$ / $$ |$$ | $$ | $$ | \\$$$$$$\\ \u001b[31m$$ | \u001b[0m $$$$$$$ |$$ | $$ | │\n", - "│ \u001b[32m$$$ / \\$$$ |\u001b[0m$$ ____|$$ |$$ | $$ |$$ | $$ | $$ |$$\\ \\____$$\\ \u001b[31m$$ | \u001b[0m$$ __$$ |$$ | $$ | │\n", - "│ \u001b[32m$$ / \\$$ |\u001b[0m\\$$$$$$$\\ $$ |\\$$$$$$$ |$$ | $$ | \\$$$$ |$$$$$$$ |\u001b[31m$$$$$$$$\\ \u001b[0m\\$$$$$$$ |$$$$$$$ | │\n", - "│ \u001b[32m\\__/ \\__|\u001b[0m \\_______|\\__| \\____$$ |\\__| \\__| \\____/ \\_______/ \u001b[31m\\________|\u001b[0m \\_______|\\_______/ │\n", - "│ \u001b[32m \u001b[0m $$\\ $$ |\u001b[31m\u001b[0m │\n", - "│ \u001b[32m \u001b[0m \\$$$$$$ |\u001b[31m\u001b[0m │\n", - "│ \u001b[32m \u001b[0m \\______/\u001b[31m\u001b[0m │\n", - "│ │\n", - "│ Inspect - Edit - Evolve Neural Networks │\n", - "│ │\n", - "╰─── By GrayBx, v1.3.3.dev7 ──────────────────────────────────────────────────────────────────────────╯\n", - "\n", - "\n", - "16/07/2026-08:32:40.394 INFO:weightslab.utils.telemetry:ping_import: WeightsLab uses anonymous usage data (package version used, and OS name). Set WL_NO_TELEMETRY=1 to disable.\n", - "Using device: cpu\n" - ] - } - ], + "outputs": [], "source": [ "import os\n", "import ssl\n", @@ -263,10 +123,6 @@ "from PIL import Image\n", "\n", "import weightslab as wl\n", - "from weightslab.components.global_monitoring import (\n", - " guard_training_context,\n", - " guard_testing_context,\n", - ")\n", "\n", "logging.basicConfig(level=logging.ERROR)\n", "logging.getLogger(\"PIL\").setLevel(logging.INFO)\n", @@ -292,7 +148,7 @@ }, { "cell_type": "code", - "execution_count": 3, + "execution_count": null, "metadata": { "id": "csaORA30_xBe" }, @@ -529,7 +385,7 @@ }, { "cell_type": "code", - "execution_count": 4, + "execution_count": null, "metadata": { "id": "GFZWMzNi_xBf" }, @@ -674,7 +530,7 @@ }, { "cell_type": "code", - "execution_count": 5, + "execution_count": null, "metadata": { "id": "GQpthcwy_xBf" }, @@ -941,7 +797,7 @@ }, { "cell_type": "code", - "execution_count": 7, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/" @@ -949,20 +805,7 @@ "id": "YZCPlGjN_xBf", "outputId": "97c94163-c64d-4f7e-8ce7-d812b5983885" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "16/07/2026-08:32:57.025 INFO:weightslab.components.checkpoint_manager:__init__: CheckpointManager initialized at /tmp/tmpf35sd4_z\n", - "16/07/2026-08:32:57.030 INFO:weightslab.src:watch_or_edit: Registered new checkpoint manager in ledger\n", - "16/07/2026-08:32:57.040 INFO:root:set_log_directory: Log file moved from /tmp/tmpmt3c2rur/weightslab_20260716_083240.log to /tmp/tmpf35sd4_z/weightslab_20260716_083240.log\n", - "16/07/2026-08:32:57.041 INFO:root:set_log_directory: Log directory updated to: /tmp/tmpf35sd4_z\n", - "16/07/2026-08:32:57.044 INFO:root:set_log_directory: Log file: /tmp/tmpf35sd4_z/weightslab_20260716_083240.log\n", - "Experiment logs -> /tmp/tmpf35sd4_z\n" - ] - } - ], + "outputs": [], "source": [ "# Inlined from the example's config.yaml (comments preserved).\n", "config = {\n", @@ -1045,7 +888,7 @@ }, { "cell_type": "code", - "execution_count": 8, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/" @@ -1053,60 +896,7 @@ "id": "Xqkj3Lo__xBg", "outputId": "ae81a09b-c057-45d3-e3b3-254d99b31e78" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "[data] Downloading Penn-Fudan dataset to ./data/PennFudanPed.zip ...\n", - "[data] Extracting Penn-Fudan ...\n", - "16/07/2026-08:33:02.465 INFO:weightslab.data.data_samples_with_ops:__init__: [DataSampleTrackingWrapper] H5 persistence enabled at /tmp/tmpf35sd4_z/checkpoints/data/data.h5\n", - "16/07/2026-08:33:02.469 INFO:weightslab.data.data_samples_with_ops:__init__: Preloading sample statistics for PHYSICAL indices in split 'train_loader' with: preload_metadata=True...\n", - "16/07/2026-08:33:02.471 INFO:weightslab.data.data_samples_with_ops:__init__: Labels will be loaded on demand from the wrapped dataset for split 'train_loader' when accessed, which may increase latency on first access but reduces initialization time.\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "Initializing ledger for split 'train_loader': 100%|██████████| 136/136 [00:00<00:00, 94114.06it/s]" - ] - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "16/07/2026-08:33:02.481 INFO:weightslab.data.dataframe_manager:register_split: [LedgeredDataFrameManager] Registering split 'train_loader' with 136 samples.\n", - "16/07/2026-08:33:02.527 INFO:weightslab.data.dataframe_manager:register_split: [LedgeredDataFrameManager] After annotation expansion: 136 samples → 136 annotation rows.\n", - "16/07/2026-08:33:02.528 INFO:weightslab.data.h5_array_store:__init__: [H5ArrayStore] Initialized with cache limit: 2048MB\n", - "16/07/2026-08:33:02.562 INFO:weightslab.data.data_samples_with_ops:__init__: [DataSampleTrackingWrapper] H5 persistence enabled at /tmp/tmpf35sd4_z/checkpoints/data/data.h5\n", - "16/07/2026-08:33:02.564 INFO:weightslab.data.data_samples_with_ops:__init__: Preloading sample statistics for PHYSICAL indices in split 'test_loader' with: preload_labels=True, preload_metadata=True...\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "\n", - "Initializing ledger for split 'test_loader': 100%|██████████| 34/34 [00:00<00:00, 102.86it/s]" - ] - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "16/07/2026-08:33:02.902 INFO:weightslab.data.dataframe_manager:register_split: [LedgeredDataFrameManager] Registering split 'test_loader' with 34 samples.\n", - "16/07/2026-08:33:02.907 INFO:weightslab.data.dataframe_manager:register_split: [LedgeredDataFrameManager] After annotation expansion: 34 samples → 34 annotation rows.\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "\n" - ] - } - ], + "outputs": [], "source": [ "# Penn-Fudan auto-downloads (~170 images) on first dataset construction - nothing to upload.\n", "data_root = config.get(\"data_root\") or \"./data\"\n", @@ -1179,7 +969,7 @@ }, { "cell_type": "code", - "execution_count": 9, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/" @@ -1187,60 +977,7 @@ "id": "21SKj-oK_xBg", "outputId": "698a493f-d5a5-4e42-c399-563b6addc2bf" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Downloading: \"https://download.pytorch.org/models/mobilenet_v3_small-047dcff4.pth\" to /root/.cache/torch/hub/checkpoints/mobilenet_v3_small-047dcff4.pth\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "100%|██████████| 9.83M/9.83M [00:00<00:00, 133MB/s]" - ] - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "16/07/2026-08:33:23.329 INFO:weightslab.backend.model_interface:__init__: Using checkpoint manager from ledger\n", - "\n", - "============================================================\n", - "Computing class weights for 1 classes (max 200 samples)...\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "\n", - " Analyzing Distribution: 100%|██████████| 136/136 [00:00<00:00, 147.35it/s]" - ] - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "\n", - "Class distribution and weights:\n", - "Class 0 (person): 100.00% -> weight: 1.000\n", - "============================================================\n", - "\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "\n", - "/tmp/ipykernel_5133/4133204388.py:171: UserWarning: To copy construct from a tensor, it is recommended to use sourceTensor.detach().clone() or sourceTensor.detach().clone().requires_grad_(True), rather than torch.tensor(sourceTensor).\n", - " \"weights\", torch.tensor(weights) if weights is not None else None\n" - ] - } - ], + "outputs": [], "source": [ "# --- Model ---\n", "_model = SmallDetector(\n", @@ -1326,7 +1063,7 @@ }, { "cell_type": "code", - "execution_count": 10, + "execution_count": null, "metadata": { "id": "g_m26pWr_xBg" }, @@ -1339,7 +1076,7 @@ " DataSampleTrackingWrapper. `targets` is per sample a [N, 6] tensor of boxes\n", " ([x1, y1, x2, y2, class_id, confidence]); see utils/data.det_collate.\n", " \"\"\"\n", - " with guard_training_context:\n", + " with wl.guard_training_context:\n", " (inputs, ids, targets, _) = next(loader)\n", " inputs = inputs.to(device)\n", " targets = [t.to(device) for t in targets]\n", @@ -1369,7 +1106,7 @@ " \"\"\"Full evaluation pass over the val loader.\"\"\"\n", " losses = 0.0\n", " ious = 0.0\n", - " with guard_testing_context, torch.no_grad():\n", + " with wl.guard_testing_context, torch.no_grad():\n", " for inputs, ids, targets, _ in loader:\n", " inputs = inputs.to(device)\n", " targets = [t.to(device) for t in targets]\n", @@ -1412,3187 +1149,7 @@ "id": "5yQOd8iv_xBg", "outputId": "19c6f02b-80b3-4eea-a207-aeef7c83d6f5" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "16/07/2026-08:33:32.039 INFO:weightslab.trainer.trainer_services:_run_security_preflight: \n", - "\n", - "# #######################################\n", - "# #######################################\n", - "[gRPC] Security preflight checks:\n", - "\tTLS: DISABLED (unencrypted traffic)\n", - "\tAuth tokens: NONE configured\n", - "\t! WARNING: GRPC_TLS_ENABLED=0. Traffic will be unencrypted. Use only for development.\n", - "\t! WARNING: No GRPC_AUTH_TOKEN/GRPC_AUTH_TOKENS configured. Only transport-level trust (TLS/mTLS) will protect RPC access.\n", - "# #######################################\n", - "# #######################################\n", - "\n", - "16/07/2026-08:33:32.041 INFO:weightslab.trainer.trainer_services:grpc_serve: [gRPC] Watchdogs disabled via WEIGHTSLAB_DISABLE_WATCHDOGS.\n", - "16/07/2026-08:33:32.045 INFO:weightslab.trainer.trainer_services:serving_thread_callback: [gRPC] Thread callback started\n", - "16/07/2026-08:33:32.045 INFO:weightslab.trainer.trainer_services:grpc_serve: [gRPC] Server started with watchdogs disabled (host=0.0.0.0 port=50051 workers=None)\n", - "16/07/2026-08:33:32.047 INFO:weightslab.trainer.trainer_services:serving_thread_callback: [gRPC] Creating ThreadPoolExecutor with 6 worker threads (n_workers_grpc=None, max_concurrent_rpcs=None)\n", - "16/07/2026-08:33:32.051 INFO:weightslab.tunnel:_ensure_bore: Downloading bore v0.6.0 (bore-v0.6.0-x86_64-unknown-linux-musl.tar.gz)...\n", - "16/07/2026-08:33:32.129 INFO:weightslab.trainer.trainer_services:serving_thread_callback: [gRPC] Server object created\n", - "16/07/2026-08:33:32.142 INFO:weightslab.trainer.services.agent.agent:__init__: Initializing DataManipulationAgent\n", - "16/07/2026-08:33:32.156 WARNING:weightslab.trainer.services.agent.agent:_load_config: Error loading config from /usr/local/lib/python3.12/dist-packages: [Errno 21] Is a directory: '/usr/local/lib/python3.12/dist-packages'\n", - "16/07/2026-08:33:32.157 INFO:weightslab.trainer.services.agent.agent:_load_config: \n", - "\n", - "# #######################################\n", - "# #######################################\n", - "Agent initialized from configuration /content/agent_config.yaml: \n", - "\tFinal Agent Configuration: Preferred Provider=openrouter, \n", - "\tFallback to Local=True, \n", - "\tOpenRouter Model=~google/gemini-flash-latest with:\n", - "\t\tAPI Key=None\n", - "\t\tBase URL=https://openrouter.ai/api/v1, \n", - "\tOllama Model=llama3.2:3b\n", - "# #######################################\n", - "# #######################################\n", - "\n", - "16/07/2026-08:33:32.158 INFO:weightslab.trainer.services.agent.agent:_setup_providers: Setting up Ollama with model llama3.2:3b\n", - "16/07/2026-08:33:32.260 INFO:weightslab.trainer.services.agent.agent:_setup_providers: [Agent] Ollama enabled: llama3.2:3b\n", - "16/07/2026-08:33:32.261 INFO:weightslab.trainer.services.data_service:_compute_natural_sort_stats: [DataService] Starting natural sort stats computation with weights: {'brightness': 0.7, 'entropy': 0.3, 'hue': 0.0}\n", - "16/07/2026-08:33:32.273 INFO:weightslab.trainer.services.data_service:_compute_natural_sort_stats: [DataService] Computing sort stats for 170 samples...\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "Computing Natural Sort Stats: 26%|██▌ | 44/170 [00:00<00:02, 44.78samples/s]" - ] - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "============================================================\n", - " Backend exposed via bore at: bore.pub:41162\n", - " On your local machine (Docker running), run:\n", - " weightslab start\n", - " weightslab tunnel bore.pub:41162\n", - "============================================================\n", - "16/07/2026-08:33:33.256 INFO:weightslab.backend.cli:cli_serve: cli_thread_started\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "\rComputing Natural Sort Stats: 31%|███ | 53/170 [00:00<00:02, 54.32samples/s]" - ] - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "16/07/2026-08:33:33.346 INFO:weightslab.src:start_training: Starting WeightsLab training mode with a timeout of 3 seconds.\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "Computing Natural Sort Stats: 100%|██████████| 170/170 [00:02<00:00, 62.04samples/s]" - ] - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "16/07/2026-08:33:35.145 INFO:weightslab.trainer.services.data_service:_compute_natural_sort_stats: [DataService] Completed stats computation for 170 samples\n", - "\n", - "\n", - "Natural sort computation finished for 170 samples\n", - "\n", - "\n", - "16/07/2026-08:33:35.147 INFO:weightslab.trainer.services.data_service:_build_preview_cache: [PreviewCache] Building 64×64 or less or less preview cache for 170 samples …\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "\n", - "/usr/local/lib/python3.12/dist-packages/weightslab/data/dataframe_manager.py:703: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise an error in a future version of pandas. Value '[7.64467454 7.63759617 7.82511253 7.60454149 7.68511861 7.70962498\n", - " 7.66256563 7.67420456 7.49258183 7.81929857 6.45717949 7.54281418\n", - " 6.79016624 7.48513405 7.32040805 7.82787458 7.77316711 7.5562755\n", - " 7.61231887 7.63537192 7.82628297 7.76547166 7.13113334 7.52482093\n", - " 7.63815989 7.7459368 7.58974044 7.75762173 7.18705444 7.48240921\n", - " 7.76600207 7.76543333 7.71791639 7.79005573]' has dtype incompatible with float32, please explicitly cast to a compatible dtype first.\n", - " self._df.loc[existing_idx, all_cols] = df_norm.loc[existing_idx, all_cols]\n", - "/usr/local/lib/python3.12/dist-packages/weightslab/data/dataframe_manager.py:703: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise an error in a future version of pandas. Value '[ 53.51681774 93.30201817 61.98435004 65.24323962 43.56923658\n", - " 93.96657441 67.03601995 73.99117037 80.06768645 84.95401979\n", - " 88.10120384 55.8291746 82.3123762 58.22935624 100.60901833\n", - " 94.05543837 79.249868 53.99237577 65.56806674 73.41094873\n", - " 99.16436384 101.13093832 93.03205657 99.69871153 101.89878035\n", - " 85.80537355 86.84879743 86.214422 99.49009009 81.44740104\n", - " 56.29921017 61.5624775 67.23737161 66.83788963]' has dtype incompatible with float32, please explicitly cast to a compatible dtype first.\n", - " self._df.loc[existing_idx, all_cols] = df_norm.loc[existing_idx, all_cols]\n", - "/usr/local/lib/python3.12/dist-packages/weightslab/data/dataframe_manager.py:703: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise an error in a future version of pandas. Value '[ 64.49066497 65.37024572 104.59268596 52.93642734 72.02098046\n", - " 48.14319779 113.24922957 46.29628522 43.88088143 25.53077158\n", - " 60.70043395 49.99961312 32.48544889 87.93221913 36.88762984\n", - " 47.55134406 57.05124766 64.13736411 65.30893805 96.26249138\n", - " 38.45269062 42.4448374 35.35399698 49.62630278 46.2604104\n", - " 42.64606061 32.72595598 68.50939554 46.54847017 59.36240918\n", - " 55.71851252 58.31763177 50.1109307 61.99646905]' has dtype incompatible with float32, please explicitly cast to a compatible dtype first.\n", - " self._df.loc[existing_idx, all_cols] = df_norm.loc[existing_idx, all_cols]\n", - "/usr/local/lib/python3.12/dist-packages/weightslab/data/dataframe_manager.py:703: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise an error in a future version of pandas. Value '[0.65953828 0.61511859 0.63585766 0.62043339 0.63605741 0.58398832\n", - " 0.53613578 0.62181845 0.62159035 0.66207224 0.38428878 0.64687088\n", - " 0.52937324 0.5747199 0.54391107 0.61459091 0.61792182 0.55884669\n", - " 0.5585666 0.58468377 0.60382653 0.58175791 0.50977316 0.59231942\n", - " 0.5993173 0.63587023 0.63463003 0.58706704 0.4947014 0.59870562\n", - " 0.58584364 0.62449927 0.61637045 0.59844273]' has dtype incompatible with float32, please explicitly cast to a compatible dtype first.\n", - " self._df.loc[existing_idx, all_cols] = df_norm.loc[existing_idx, all_cols]\n" - ] - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "16/07/2026-08:33:35.148 INFO:weightslab.trainer.services.data_service:__init__: DataService initialized.\n", - "16/07/2026-08:33:35.150 INFO:weightslab.trainer.trainer_services:serving_thread_callback: [gRPC] Servicer added\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "\r[PreviewCache]: 0%| | 0/170 [00:00 16 (Loader: train_loader)\n", - "16/07/2026-08:40:17.818 INFO:weightslab.components.global_monitoring:resume: \n", - "Training resumed as modules hashes have been computed: ['bfb454d6', 'd39e7bd9', '43e44972'].\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "/usr/local/lib/python3.12/dist-packages/torch/utils/data/dataloader.py:668: UserWarning: 'pin_memory' argument is set as true but no accelerator is found, then device pinned memory won't be used.\n", - " warnings.warn(warn_msg)\n", - "\n", - "Training: 20%|█▉ | 296/1500 [06:41<33:32, 1.67s/it, iou=56.44%, test_loss=8.1024, train_loss=1.8979]\u001b[A\n", - "Training: 20%|█▉ | 297/1500 [06:41<10:23:18, 31.09s/it, iou=56.44%, test_loss=8.1024, train_loss=1.8979]\u001b[A\n", - "Training: 20%|█▉ | 297/1500 [06:41<10:23:18, 31.09s/it, iou=56.44%, test_loss=8.1024, train_loss=2.4030]\u001b[A\n", - "Training: 20%|█▉ | 298/1500 [06:41<7:18:34, 21.89s/it, iou=56.44%, test_loss=8.1024, train_loss=2.4030] \u001b[A\n", - "Training: 20%|█▉ | 298/1500 [06:42<7:18:34, 21.89s/it, iou=56.44%, test_loss=8.1024, train_loss=1.9995]\u001b[A\n", - "Training: 20%|█▉ | 299/1500 [06:42<5:11:45, 15.57s/it, iou=56.44%, test_loss=8.1024, train_loss=1.9995]\u001b[A\n", - "Training: 20%|█▉ | 299/1500 [06:42<5:11:45, 15.57s/it, iou=56.44%, test_loss=8.1024, train_loss=1.9483]\u001b[A/usr/local/lib/python3.12/dist-packages/torch/utils/data/dataloader.py:668: UserWarning: 'pin_memory' argument is set as true but no accelerator is found, then device pinned memory won't be used.\n", - " warnings.warn(warn_msg)\n" - ] - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "16/07/2026-08:40:20.402 INFO:weightslab.components.checkpoint_manager:save_model_checkpoint: Saved model checkpoint: bfb454d6d39e7bd943e44972_step_000300.pt\n", - "16/07/2026-08:40:20.588 INFO:weightslab.components.checkpoint_manager:save_logger_snapshot: Saved logger snapshot: /tmp/tmpf35sd4_z/checkpoints/loggers/loggers.manifest.json (1 chunks)\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "\n", - "Training: 20%|█▉ | 299/1500 [06:43<5:11:45, 15.57s/it, iou=56.44%, test_loss=8.1024, train_loss=2.3940]\u001b[A\n", - "Training: 20%|██ | 301/1500 [06:43<2:51:19, 8.57s/it, iou=56.44%, test_loss=8.1024, train_loss=2.3940]\u001b[A" - ] - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "16/07/2026-08:40:21.165 INFO:weightslab.backend.dataloader_interface:set_batch_size: Batch size updated: 8 -> 16 (Loader: test_loader)\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "/usr/local/lib/python3.12/dist-packages/torch/utils/data/dataloader.py:668: UserWarning: 'pin_memory' argument is set as true but no accelerator is found, then device pinned memory won't be used.\n", - " warnings.warn(warn_msg)\n" - ] - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "16/07/2026-08:40:23.499 INFO:weightslab.src:write_dataframe: write_dataframe: wrote 593 row(s) × 24 column(s) as json to /tmp/tmpf35sd4_z/d8e01c87_dataframe.json\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "\n", - "Training: 20%|██ | 301/1500 [06:46<2:51:19, 8.57s/it, iou=33.90%, test_loss=4.7464, train_loss=1.9816]\u001b[A\n", - "Training: 20%|██ | 302/1500 [06:46<2:22:34, 7.14s/it, iou=33.90%, test_loss=4.7464, train_loss=1.9816]\u001b[A\n", - "Training: 20%|██ | 302/1500 [06:46<2:22:34, 7.14s/it, iou=33.90%, test_loss=4.7464, train_loss=1.8739]\u001b[A\n", - "Training: 20%|██ | 303/1500 [06:46<1:47:37, 5.39s/it, iou=33.90%, test_loss=4.7464, train_loss=1.8739]\u001b[A\n", - "Training: 20%|██ | 303/1500 [06:46<1:47:37, 5.39s/it, iou=33.90%, test_loss=4.7464, train_loss=1.8943]\u001b[A/usr/local/lib/python3.12/dist-packages/torch/utils/data/dataloader.py:668: UserWarning: 'pin_memory' argument is set as true but no accelerator is found, then device pinned memory won't be used.\n", - " warnings.warn(warn_msg)\n", - "\n", - "Training: 20%|██ | 303/1500 [06:47<1:47:37, 5.39s/it, iou=33.90%, test_loss=4.7464, train_loss=1.8743]\u001b[A\n", - "Training: 20%|██ | 305/1500 [06:47<1:03:34, 3.19s/it, iou=33.90%, test_loss=4.7464, train_loss=1.8743]\u001b[A\n", - "Training: 20%|██ | 305/1500 [06:47<1:03:34, 3.19s/it, iou=33.90%, test_loss=4.7464, train_loss=1.9137]\u001b[A\n", - "Training: 20%|██ | 306/1500 [06:47<50:33, 2.54s/it, iou=33.90%, test_loss=4.7464, train_loss=1.9137] \u001b[A\n", - "Training: 20%|██ | 306/1500 [06:48<50:33, 2.54s/it, iou=33.90%, test_loss=4.7464, train_loss=2.4774]\u001b[A\n", - "Training: 20%|██ | 307/1500 [06:48<42:26, 2.13s/it, iou=33.90%, test_loss=4.7464, train_loss=2.4774]\u001b[A\n", - "Training: 20%|██ | 307/1500 [06:48<42:26, 2.13s/it, iou=33.90%, test_loss=4.7464, train_loss=1.8858]\u001b[A/usr/local/lib/python3.12/dist-packages/torch/utils/data/dataloader.py:668: UserWarning: 'pin_memory' argument is set as true but no accelerator is found, then device pinned memory won't be used.\n", - " warnings.warn(warn_msg)\n", - "\n", - "Training: 20%|██ | 307/1500 [06:49<42:26, 2.13s/it, iou=33.90%, test_loss=4.7464, train_loss=1.9598]\u001b[A\n", - "Training: 21%|██ | 309/1500 [06:49<26:52, 1.35s/it, iou=33.90%, test_loss=4.7464, train_loss=1.9598]\u001b[A\n", - "Training: 21%|██ | 309/1500 [06:49<26:52, 1.35s/it, iou=33.90%, test_loss=4.7464, train_loss=2.3849]\u001b[A\n", - "Training: 21%|██ | 310/1500 [06:49<22:41, 1.14s/it, iou=33.90%, test_loss=4.7464, train_loss=2.3849]\u001b[A\n", - "Training: 21%|██ | 310/1500 [06:50<22:41, 1.14s/it, iou=33.90%, test_loss=4.7464, train_loss=1.8953]\u001b[A\n", - "Training: 21%|██ | 311/1500 [06:50<21:04, 1.06s/it, iou=33.90%, test_loss=4.7464, train_loss=1.8953]\u001b[A\n", - "Training: 21%|██ | 311/1500 [06:50<21:04, 1.06s/it, iou=33.90%, test_loss=4.7464, train_loss=1.9958]\u001b[A\n", - "Training: 21%|██ | 312/1500 [06:50<16:32, 1.20it/s, iou=33.90%, test_loss=4.7464, train_loss=1.9958]\u001b[A/usr/local/lib/python3.12/dist-packages/torch/utils/data/dataloader.py:668: UserWarning: 'pin_memory' argument is set as true but no accelerator is found, then device pinned memory won't be used.\n", - " warnings.warn(warn_msg)\n", - "\n", - "Training: 21%|██ | 312/1500 [06:51<16:32, 1.20it/s, iou=33.90%, test_loss=4.7464, train_loss=1.8536]\u001b[A\n", - "Training: 21%|██ | 313/1500 [06:51<14:17, 1.38it/s, iou=33.90%, test_loss=4.7464, train_loss=1.8536]\u001b[A\n", - "Training: 21%|██ | 313/1500 [06:51<14:17, 1.38it/s, iou=33.90%, test_loss=4.7464, train_loss=2.3955]\u001b[A\n", - "Training: 21%|██ | 314/1500 [06:51<13:25, 1.47it/s, iou=33.90%, test_loss=4.7464, train_loss=2.3955]\u001b[A\n", - "Training: 21%|██ | 314/1500 [06:52<13:25, 1.47it/s, iou=33.90%, test_loss=4.7464, train_loss=1.9691]\u001b[A\n", - "Training: 21%|██ | 315/1500 [06:52<16:04, 1.23it/s, iou=33.90%, test_loss=4.7464, train_loss=1.9691]\u001b[A\n", - "Training: 21%|██ | 315/1500 [06:53<16:04, 1.23it/s, iou=33.90%, test_loss=4.7464, train_loss=1.8806]\u001b[A\n", - "Training: 21%|██ | 316/1500 [06:53<13:43, 1.44it/s, iou=33.90%, test_loss=4.7464, train_loss=1.8806]\u001b[A/usr/local/lib/python3.12/dist-packages/torch/utils/data/dataloader.py:668: UserWarning: 'pin_memory' argument is set as true but no accelerator is found, then device pinned memory won't be used.\n", - " warnings.warn(warn_msg)\n", - "\n", - "Training: 21%|██ | 316/1500 [06:54<13:43, 1.44it/s, iou=33.90%, test_loss=4.7464, train_loss=1.8700]\u001b[A\n", - "Training: 21%|██ | 317/1500 [06:54<17:55, 1.10it/s, iou=33.90%, test_loss=4.7464, train_loss=1.8700]\u001b[A\n", - "Training: 21%|██ | 317/1500 [06:55<17:55, 1.10it/s, iou=33.90%, test_loss=4.7464, train_loss=1.9674]\u001b[A\n", - "Training: 21%|██ | 318/1500 [06:55<16:26, 1.20it/s, iou=33.90%, test_loss=4.7464, train_loss=1.9674]\u001b[A\n", - "Training: 21%|██ | 318/1500 [06:56<16:26, 1.20it/s, iou=33.90%, test_loss=4.7464, train_loss=2.3859]\u001b[A\n", - "Training: 21%|██▏ | 319/1500 [06:56<18:06, 1.09it/s, iou=33.90%, test_loss=4.7464, train_loss=2.3859]\u001b[A\n", - "Training: 21%|██▏ | 319/1500 [06:56<18:06, 1.09it/s, iou=33.90%, test_loss=4.7464, train_loss=1.8463]\u001b[A/usr/local/lib/python3.12/dist-packages/torch/utils/data/dataloader.py:668: UserWarning: 'pin_memory' argument is set as true but no accelerator is found, then device pinned memory won't be used.\n", - " warnings.warn(warn_msg)\n", - "\n", - "Training: 21%|██▏ | 319/1500 [06:56<18:06, 1.09it/s, iou=33.90%, test_loss=4.7464, train_loss=1.9668]\u001b[A\n", - "Training: 21%|██▏ | 321/1500 [06:56<11:56, 1.64it/s, iou=33.90%, test_loss=4.7464, train_loss=1.9668]\u001b[A\n", - "Training: 21%|██▏ | 321/1500 [06:57<11:56, 1.64it/s, iou=33.90%, test_loss=4.7464, train_loss=1.8672]\u001b[A\n", - "Training: 21%|██▏ | 322/1500 [06:57<11:02, 1.78it/s, iou=33.90%, test_loss=4.7464, train_loss=1.8672]\u001b[A\n", - "Training: 21%|██▏ | 322/1500 [06:58<11:02, 1.78it/s, iou=33.90%, test_loss=4.7464, train_loss=2.3978]\u001b[A\n", - "Training: 22%|██▏ | 323/1500 [06:58<12:20, 1.59it/s, iou=33.90%, test_loss=4.7464, train_loss=2.3978]\u001b[A\n", - "Training: 22%|██▏ | 323/1500 [06:58<12:20, 1.59it/s, iou=33.90%, test_loss=4.7464, train_loss=1.9003]\u001b[A\n", - "Training: 22%|██▏ | 324/1500 [06:58<09:58, 1.96it/s, iou=33.90%, test_loss=4.7464, train_loss=1.9003]\u001b[A/usr/local/lib/python3.12/dist-packages/torch/utils/data/dataloader.py:668: UserWarning: 'pin_memory' argument is set as true but no accelerator is found, then device pinned memory won't be used.\n", - " warnings.warn(warn_msg)\n", - "\n", - "Training: 22%|██▏ | 324/1500 [06:58<09:58, 1.96it/s, iou=33.90%, test_loss=4.7464, train_loss=1.8884]\u001b[A\n", - "Training: 22%|██▏ | 325/1500 [06:58<09:27, 2.07it/s, iou=33.90%, test_loss=4.7464, train_loss=1.8884]\u001b[A" - ] - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "16/07/2026-08:40:36.225 INFO:weightslab.components.checkpoint_manager:save_model_checkpoint: Saved model checkpoint: bfb454d6d39e7bd943e44972_step_000325.pt\n", - "16/07/2026-08:40:36.404 INFO:weightslab.components.checkpoint_manager:save_logger_snapshot: Saved logger snapshot: /tmp/tmpf35sd4_z/checkpoints/loggers/loggers.manifest.json (1 chunks)\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "\n", - "Training: 22%|██▏ | 325/1500 [06:59<09:27, 2.07it/s, iou=33.90%, test_loss=4.7464, train_loss=1.9650]\u001b[A\n", - "Training: 22%|██▏ | 326/1500 [06:59<10:33, 1.85it/s, iou=33.90%, test_loss=4.7464, train_loss=1.9650]\u001b[A/usr/local/lib/python3.12/dist-packages/torch/utils/data/dataloader.py:668: UserWarning: 'pin_memory' argument is set as true but no accelerator is found, then device pinned memory won't be used.\n", - " warnings.warn(warn_msg)\n" - ] - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "16/07/2026-08:40:38.855 INFO:weightslab.components.global_monitoring:pause: \n", - "Training paused.\n", - "16/07/2026-08:40:38.862 INFO:weightslab.trainer.services.experiment_service:ExperimentCommand: [WeightsLab] UI Command: HP changed, experiment paused!\n", - "16/07/2026-08:40:38.865 INFO:weightslab.trainer.services.experiment_service:ExperimentCommand: \n", - "[WeightsLab] UI Command: PAUSE\n", - "16/07/2026-08:40:38.870 INFO:weightslab.components.global_monitoring:pause: \n", - "Training paused.\n", - "16/07/2026-08:40:39.412 INFO:weightslab.src:write_dataframe: write_dataframe: wrote 593 row(s) × 24 column(s) as json to /tmp/tmpf35sd4_z/d8e01c87_dataframe.json\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "\n", - "Training: 22%|██▏ | 326/1500 [07:02<10:33, 1.85it/s, iou=54.03%, test_loss=7.9861, train_loss=2.4054]\u001b[A\n", - "Training: 22%|██▏ | 327/1500 [07:02<23:56, 1.22s/it, iou=54.03%, test_loss=7.9861, train_loss=2.4054]\u001b[A" - ] - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "16/07/2026-08:40:41.095 INFO:weightslab.trainer.services.experiment_service:RestoreCheckpoint: Restoring checkpoint from hash: 802319f761e1ba3dc3472a78 (weights-only, target_step=24)\n", - "16/07/2026-08:40:41.096 INFO:weightslab.trainer.services.experiment_service:RestoreCheckpoint: Pausing training before restore...\n", - "16/07/2026-08:40:41.099 INFO:weightslab.components.global_monitoring:pause: \n", - "Training paused.\n", - "16/07/2026-08:40:41.102 INFO:weightslab.components.checkpoint_manager:load_state: \n", - "============================================================\n", - "16/07/2026-08:40:41.103 INFO:weightslab.components.checkpoint_manager:load_state: Loading and applying state: 802319f761e1ba3d...\n", - "16/07/2026-08:40:41.104 INFO:weightslab.components.checkpoint_manager:load_state: ============================================================\n", - "16/07/2026-08:40:41.110 INFO:weightslab.components.checkpoint_manager:load_checkpoint: Loading checkpoint 802319f761e1ba3d...\n", - "16/07/2026-08:40:41.112 INFO:weightslab.components.checkpoint_manager:load_checkpoint: Target: HP=802319f7 MODEL=61e1ba3d DATA=c3472a78\n", - "16/07/2026-08:40:41.115 INFO:weightslab.components.checkpoint_manager:load_checkpoint: Current: HP=bfb454d6 MODEL=d39e7bd9 DATA=43e44972\n", - "16/07/2026-08:40:41.118 WARNING:weightslab.components.checkpoint_manager:load_checkpoint: [WARNING] Model architecture file not found: /tmp/tmpf35sd4_z/checkpoints/models/61e1ba3d/61e1ba3d_architecture.pkl\n", - "16/07/2026-08:40:41.176 INFO:weightslab.components.checkpoint_manager:load_checkpoint: [OK] Loaded weights from step 25 with RNG state\n", - "16/07/2026-08:40:41.184 INFO:weightslab.components.checkpoint_manager:load_checkpoint: [OK] Loaded config (hash changed)\n", - "16/07/2026-08:40:41.198 INFO:weightslab.components.checkpoint_manager:load_checkpoint: [OK] Loaded data snapshot (170 rows) with RNG state\n", - "16/07/2026-08:40:41.200 INFO:weightslab.components.checkpoint_manager:load_checkpoint: Loaded components: {'data', 'weights', 'config'}\n", - "16/07/2026-08:40:41.215 INFO:weightslab.components.checkpoint_manager:load_state: [OK] Applied weights to existing model (step 25)\n", - "16/07/2026-08:40:41.217 INFO:weightslab.components.checkpoint_manager:load_state: Loaded optimizer state\n", - "16/07/2026-08:40:41.219 INFO:weightslab.components.checkpoint_manager:load_state: [OK] Applied hyperparameters config\n", - "16/07/2026-08:40:41.254 INFO:weightslab.components.checkpoint_manager:load_state: [OK] Applied data snapshot (170 rows)\n", - "16/07/2026-08:40:41.266 INFO:weightslab.components.checkpoint_manager:load_state: \n", - "[OK] Successfully loaded and applied state: 802319f761e1ba3d\n", - "16/07/2026-08:40:41.267 INFO:weightslab.components.checkpoint_manager:load_state: ============================================================\n", - "\n", - "16/07/2026-08:40:41.272 INFO:weightslab.trainer.services.experiment_service:RestoreCheckpoint: Successfully restored checkpoint: 802319f761e1ba3dc3472a78\n", - "16/07/2026-08:40:44.091 INFO:weightslab.components.global_monitoring:pause: \n", - "Training paused.\n", - "16/07/2026-08:40:44.092 INFO:weightslab.trainer.services.experiment_service:ExperimentCommand: [WeightsLab] UI Command: HP changed, experiment paused!\n", - "16/07/2026-08:40:44.093 INFO:weightslab.trainer.services.experiment_service:ExperimentCommand: \n", - "[WeightsLab] UI Command: RESUME\n", - "16/07/2026-08:40:44.096 INFO:weightslab.components.global_monitoring:resume: \n", - "Attempting to resume training...\n", - "16/07/2026-08:40:44.104 INFO:weightslab.components.experiment_hash:has_changed: Experiment configuration changed: {'model'}\n", - "16/07/2026-08:40:44.107 INFO:weightslab.components.experiment_hash:generate_hash: Generated experiment hash: 802319f7f5ee131bc3472a78- (HP: 802319f7, Model: f5ee131b, Data: c3472a78)\n", - "16/07/2026-08:40:44.111 INFO:weightslab.components.checkpoint_manager:update_experiment_hash: New experiment hash: 802319f7-f5ee131b-c3472a78 (previous: 802319f7-61e1ba3d-c3472a78)\n", - "16/07/2026-08:40:44.113 INFO:weightslab.components.checkpoint_manager:update_experiment_hash: Changed components: {'model'}\n", - "16/07/2026-08:40:44.115 INFO:weightslab.components.checkpoint_manager:update_experiment_hash: Changes pending (not dumped yet): {'model'}\n", - "16/07/2026-08:40:44.118 INFO:weightslab.components.checkpoint_manager:save_pending_changes: Dumping pending changes: {'model'}\n", - "16/07/2026-08:40:44.427 INFO:weightslab.components.checkpoint_manager:save_logger_snapshot: Saved logger snapshot: /tmp/tmpf35sd4_z/checkpoints/loggers/loggers.manifest.json (1 chunks)\n", - "16/07/2026-08:40:44.434 INFO:weightslab.components.global_monitoring:resume: Hashes by module: ['802319f7', 'f5ee131b', 'c3472a78']\n", - "16/07/2026-08:40:44.437 INFO:weightslab.components.global_monitoring:resume: Resuming training now...\n", - "16/07/2026-08:40:44.438 INFO:weightslab.components.global_monitoring:resume: Hashes by module on resume: ['802319f7', 'f5ee131b', 'c3472a78']\n", - "16/07/2026-08:40:44.448 INFO:weightslab.backend.dataloader_interface:set_batch_size: Batch size updated: 16 -> 8 (Loader: train_loader)\n", - "16/07/2026-08:40:44.452 INFO:weightslab.components.global_monitoring:resume: \n", - "Training resumed as modules hashes have been computed: ['802319f7', 'f5ee131b', 'c3472a78'].\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "/usr/local/lib/python3.12/dist-packages/torch/utils/data/dataloader.py:668: UserWarning: 'pin_memory' argument is set as true but no accelerator is found, then device pinned memory won't be used.\n", - " warnings.warn(warn_msg)\n", - "\n", - "Training: 22%|██▏ | 327/1500 [07:07<23:56, 1.22s/it, iou=54.03%, test_loss=7.9861, train_loss=10.4120]\u001b[A\n", - "Training: 22%|██▏ | 328/1500 [07:07<48:08, 2.46s/it, iou=54.03%, test_loss=7.9861, train_loss=10.4120]\u001b[A\n", - "Training: 22%|██▏ | 328/1500 [07:08<48:08, 2.46s/it, iou=54.03%, test_loss=7.9861, train_loss=8.1298] \u001b[A\n", - "Training: 22%|██▏ | 329/1500 [07:08<36:27, 1.87s/it, iou=54.03%, test_loss=7.9861, train_loss=8.1298]\u001b[A\n", - "Training: 22%|██▏ | 329/1500 [07:08<36:27, 1.87s/it, iou=54.03%, test_loss=7.9861, train_loss=8.6498]\u001b[A\n", - "Training: 22%|██▏ | 330/1500 [07:08<27:35, 1.41s/it, iou=54.03%, test_loss=7.9861, train_loss=8.6498]\u001b[A\n", - "Training: 22%|██▏ | 330/1500 [07:09<27:35, 1.41s/it, iou=54.03%, test_loss=7.9861, train_loss=8.4801]\u001b[A\n", - "Training: 22%|██▏ | 331/1500 [07:09<22:37, 1.16s/it, iou=54.03%, test_loss=7.9861, train_loss=8.4801]\u001b[A" - ] - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "16/07/2026-08:40:48.449 INFO:weightslab.components.checkpoint_manager:save_model_checkpoint: Saved model checkpoint: 802319f7f5ee131bc3472a78_step_000030.pt\n", - "16/07/2026-08:40:48.825 INFO:weightslab.components.checkpoint_manager:save_logger_snapshot: Saved logger snapshot: /tmp/tmpf35sd4_z/checkpoints/loggers/loggers.manifest.json (1 chunks)\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "\n", - "Training: 22%|██▏ | 331/1500 [07:11<22:37, 1.16s/it, iou=54.03%, test_loss=7.9861, train_loss=9.8237]\u001b[A\n", - "Training: 22%|██▏ | 332/1500 [07:11<31:52, 1.64s/it, iou=54.03%, test_loss=7.9861, train_loss=9.8237]\u001b[A" - ] - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "16/07/2026-08:40:51.421 INFO:weightslab.backend.dataloader_interface:set_batch_size: Batch size updated: 16 -> 8 (Loader: test_loader)\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "/usr/local/lib/python3.12/dist-packages/torch/utils/data/dataloader.py:668: UserWarning: 'pin_memory' argument is set as true but no accelerator is found, then device pinned memory won't be used.\n", - " warnings.warn(warn_msg)\n" - ] - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "16/07/2026-08:40:56.551 WARNING:weightslab.watchdog.grpc_watchdog:unary_unary_wrapper: [gRPC] /ExperimentService/GetDataSamples completed in 3558.8ms\n", - "16/07/2026-08:41:00.986 INFO:weightslab.src:write_dataframe: write_dataframe: wrote 593 row(s) × 24 column(s) as json to /tmp/tmpf35sd4_z/d8e01c87_dataframe.json\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "\n", - "Training: 22%|██▏ | 332/1500 [07:23<31:52, 1.64s/it, iou=90.43%, test_loss=15.9695, train_loss=9.6350]\u001b[A\n", - "Training: 22%|██▏ | 333/1500 [07:23<1:32:06, 4.74s/it, iou=90.43%, test_loss=15.9695, train_loss=9.6350]\u001b[A\n", - "Training: 22%|██▏ | 333/1500 [07:24<1:32:06, 4.74s/it, iou=90.43%, test_loss=15.9695, train_loss=7.9581]\u001b[A\n", - "Training: 22%|██▏ | 334/1500 [07:24<1:06:21, 3.41s/it, iou=90.43%, test_loss=15.9695, train_loss=7.9581]\u001b[A\n", - "Training: 22%|██▏ | 334/1500 [07:24<1:06:21, 3.41s/it, iou=90.43%, test_loss=15.9695, train_loss=7.3944]\u001b[A\n", - "Training: 22%|██▏ | 335/1500 [07:24<47:55, 2.47s/it, iou=90.43%, test_loss=15.9695, train_loss=7.3944] \u001b[A\n", - "Training: 22%|██▏ | 335/1500 [07:24<47:55, 2.47s/it, iou=90.43%, test_loss=15.9695, train_loss=10.4550]\u001b[A\n", - "Training: 22%|██▏ | 336/1500 [07:24<35:17, 1.82s/it, iou=90.43%, test_loss=15.9695, train_loss=10.4550]\u001b[A" - ] - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "16/07/2026-08:41:02.160 INFO:weightslab.components.checkpoint_manager:save_model_checkpoint: Saved model checkpoint: 802319f7f5ee131bc3472a78_step_000035.pt\n", - "16/07/2026-08:41:02.346 INFO:weightslab.components.checkpoint_manager:save_logger_snapshot: Saved logger snapshot: /tmp/tmpf35sd4_z/checkpoints/loggers/loggers.manifest.json (1 chunks)\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "\n", - "Training: 22%|██▏ | 336/1500 [07:25<35:17, 1.82s/it, iou=90.43%, test_loss=15.9695, train_loss=7.6845] \u001b[A\n", - "Training: 22%|██▏ | 337/1500 [07:25<29:14, 1.51s/it, iou=90.43%, test_loss=15.9695, train_loss=7.6845]\u001b[A/usr/local/lib/python3.12/dist-packages/torch/utils/data/dataloader.py:668: UserWarning: 'pin_memory' argument is set as true but no accelerator is found, then device pinned memory won't be used.\n", - " warnings.warn(warn_msg)\n" - ] - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "16/07/2026-08:41:05.682 INFO:weightslab.src:write_dataframe: write_dataframe: wrote 593 row(s) × 24 column(s) as json to /tmp/tmpf35sd4_z/d8e01c87_dataframe.json\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "\n", - "Training: 22%|██▏ | 337/1500 [07:28<29:14, 1.51s/it, iou=55.45%, test_loss=9.3424, train_loss=6.8015] \u001b[A\n", - "Training: 23%|██▎ | 338/1500 [07:28<38:06, 1.97s/it, iou=55.45%, test_loss=9.3424, train_loss=6.8015]\u001b[A\n", - "Training: 23%|██▎ | 338/1500 [07:28<38:06, 1.97s/it, iou=55.45%, test_loss=9.3424, train_loss=8.4716]\u001b[A\n", - "Training: 23%|██▎ | 339/1500 [07:28<28:11, 1.46s/it, iou=55.45%, test_loss=9.3424, train_loss=8.4716]\u001b[A\n", - "Training: 23%|██▎ | 339/1500 [07:29<28:11, 1.46s/it, iou=55.45%, test_loss=9.3424, train_loss=8.7423]\u001b[A\n", - "Training: 23%|██▎ | 340/1500 [07:29<21:16, 1.10s/it, iou=55.45%, test_loss=9.3424, train_loss=8.7423]\u001b[A\n", - "Training: 23%|██▎ | 340/1500 [07:29<21:16, 1.10s/it, iou=55.45%, test_loss=9.3424, train_loss=8.9636]\u001b[A\n", - "Training: 23%|██▎ | 341/1500 [07:29<16:33, 1.17it/s, iou=55.45%, test_loss=9.3424, train_loss=8.9636]\u001b[A" - ] - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "16/07/2026-08:41:06.887 INFO:weightslab.components.checkpoint_manager:save_model_checkpoint: Saved model checkpoint: 802319f7f5ee131bc3472a78_step_000040.pt\n", - "16/07/2026-08:41:07.490 INFO:weightslab.components.checkpoint_manager:save_logger_snapshot: Saved logger snapshot: /tmp/tmpf35sd4_z/checkpoints/loggers/loggers.manifest.json (1 chunks)\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "\n", - "Training: 23%|██▎ | 341/1500 [07:30<16:33, 1.17it/s, iou=55.45%, test_loss=9.3424, train_loss=8.2650]\u001b[A\n", - "Training: 23%|██▎ | 342/1500 [07:30<17:36, 1.10it/s, iou=55.45%, test_loss=9.3424, train_loss=8.2650]\u001b[A/usr/local/lib/python3.12/dist-packages/torch/utils/data/dataloader.py:668: UserWarning: 'pin_memory' argument is set as true but no accelerator is found, then device pinned memory won't be used.\n", - " warnings.warn(warn_msg)\n" - ] - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "16/07/2026-08:41:10.072 INFO:weightslab.src:write_dataframe: write_dataframe: wrote 593 row(s) × 24 column(s) as json to /tmp/tmpf35sd4_z/d8e01c87_dataframe.json\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "\n", - "Training: 23%|██▎ | 342/1500 [07:32<17:36, 1.10it/s, iou=57.76%, test_loss=8.9165, train_loss=9.8330]\u001b[A\n", - "Training: 23%|██▎ | 343/1500 [07:32<26:56, 1.40s/it, iou=57.76%, test_loss=8.9165, train_loss=9.8330]\u001b[A\n", - "Training: 23%|██▎ | 343/1500 [07:33<26:56, 1.40s/it, iou=57.76%, test_loss=8.9165, train_loss=8.5029]\u001b[A\n", - "Training: 23%|██▎ | 344/1500 [07:33<20:40, 1.07s/it, iou=57.76%, test_loss=8.9165, train_loss=8.5029]\u001b[A/usr/local/lib/python3.12/dist-packages/torch/utils/data/dataloader.py:668: UserWarning: 'pin_memory' argument is set as true but no accelerator is found, then device pinned memory won't be used.\n", - " warnings.warn(warn_msg)\n", - "\n", - "Training: 23%|██▎ | 344/1500 [07:33<20:40, 1.07s/it, iou=57.76%, test_loss=8.9165, train_loss=6.6450]\u001b[A\n", - "Training: 23%|██▎ | 345/1500 [07:33<16:02, 1.20it/s, iou=57.76%, test_loss=8.9165, train_loss=6.6450]\u001b[A\n", - "Training: 23%|██▎ | 345/1500 [07:33<16:02, 1.20it/s, iou=57.76%, test_loss=8.9165, train_loss=7.5255]\u001b[A\n", - "Training: 23%|██▎ | 346/1500 [07:33<13:05, 1.47it/s, iou=57.76%, test_loss=8.9165, train_loss=7.5255]\u001b[A" - ] - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "16/07/2026-08:41:11.459 INFO:weightslab.components.checkpoint_manager:save_model_checkpoint: Saved model checkpoint: 802319f7f5ee131bc3472a78_step_000045.pt\n", - "16/07/2026-08:41:11.769 INFO:weightslab.components.checkpoint_manager:save_logger_snapshot: Saved logger snapshot: /tmp/tmpf35sd4_z/checkpoints/loggers/loggers.manifest.json (1 chunks)\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "\n", - "Training: 23%|██▎ | 346/1500 [07:34<13:05, 1.47it/s, iou=57.76%, test_loss=8.9165, train_loss=7.0646]\u001b[A\n", - "Training: 23%|██▎ | 347/1500 [07:34<15:19, 1.25it/s, iou=57.76%, test_loss=8.9165, train_loss=7.0646]\u001b[A/usr/local/lib/python3.12/dist-packages/torch/utils/data/dataloader.py:668: UserWarning: 'pin_memory' argument is set as true but no accelerator is found, then device pinned memory won't be used.\n", - " warnings.warn(warn_msg)\n" - ] - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "16/07/2026-08:41:16.032 INFO:weightslab.src:write_dataframe: write_dataframe: wrote 593 row(s) × 24 column(s) as json to /tmp/tmpf35sd4_z/d8e01c87_dataframe.json\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "\n", - "Training: 23%|██▎ | 347/1500 [07:38<15:19, 1.25it/s, iou=58.16%, test_loss=8.7419, train_loss=8.9887]\u001b[A\n", - "Training: 23%|██▎ | 348/1500 [07:38<33:36, 1.75s/it, iou=58.16%, test_loss=8.7419, train_loss=8.9887]\u001b[A\n", - "Training: 23%|██▎ | 348/1500 [07:39<33:36, 1.75s/it, iou=58.16%, test_loss=8.7419, train_loss=7.1660]\u001b[A\n", - "Training: 23%|██▎ | 349/1500 [07:39<25:01, 1.30s/it, iou=58.16%, test_loss=8.7419, train_loss=7.1660]\u001b[A\n", - "Training: 23%|██▎ | 349/1500 [07:39<25:01, 1.30s/it, iou=58.16%, test_loss=8.7419, train_loss=7.0642]\u001b[A\n", - "Training: 23%|██▎ | 350/1500 [07:39<18:58, 1.01it/s, iou=58.16%, test_loss=8.7419, train_loss=7.0642]\u001b[A\n", - "Training: 23%|██▎ | 350/1500 [07:39<18:58, 1.01it/s, iou=58.16%, test_loss=8.7419, train_loss=5.4376]\u001b[A\n", - "Training: 23%|██▎ | 351/1500 [07:39<14:41, 1.30it/s, iou=58.16%, test_loss=8.7419, train_loss=5.4376]\u001b[A" - ] - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "16/07/2026-08:41:17.068 INFO:weightslab.components.checkpoint_manager:save_model_checkpoint: Saved model checkpoint: 802319f7f5ee131bc3472a78_step_000050.pt\n", - "16/07/2026-08:41:17.255 INFO:weightslab.components.checkpoint_manager:save_logger_snapshot: Saved logger snapshot: /tmp/tmpf35sd4_z/checkpoints/loggers/loggers.manifest.json (1 chunks)\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "\n", - "Training: 23%|██▎ | 351/1500 [07:40<14:41, 1.30it/s, iou=58.16%, test_loss=8.7419, train_loss=8.3300]\u001b[A\n", - "Training: 23%|██▎ | 352/1500 [07:40<14:17, 1.34it/s, iou=58.16%, test_loss=8.7419, train_loss=8.3300]\u001b[A/usr/local/lib/python3.12/dist-packages/torch/utils/data/dataloader.py:668: UserWarning: 'pin_memory' argument is set as true but no accelerator is found, then device pinned memory won't be used.\n", - " warnings.warn(warn_msg)\n" - ] - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "16/07/2026-08:41:20.421 INFO:weightslab.src:write_dataframe: write_dataframe: wrote 593 row(s) × 24 column(s) as json to /tmp/tmpf35sd4_z/d8e01c87_dataframe.json\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "\n", - "Training: 23%|██▎ | 352/1500 [07:43<14:17, 1.34it/s, iou=58.39%, test_loss=8.7055, train_loss=6.9667]\u001b[A\n", - "Training: 24%|██▎ | 353/1500 [07:43<26:44, 1.40s/it, iou=58.39%, test_loss=8.7055, train_loss=6.9667]\u001b[A\n", - "Training: 24%|██▎ | 353/1500 [07:43<26:44, 1.40s/it, iou=58.39%, test_loss=8.7055, train_loss=5.6257]\u001b[A\n", - "Training: 24%|██▎ | 354/1500 [07:43<20:19, 1.06s/it, iou=58.39%, test_loss=8.7055, train_loss=5.6257]\u001b[A\n", - "Training: 24%|██▎ | 354/1500 [07:43<20:19, 1.06s/it, iou=58.39%, test_loss=8.7055, train_loss=6.9397]\u001b[A\n", - "Training: 24%|██▎ | 355/1500 [07:43<15:46, 1.21it/s, iou=58.39%, test_loss=8.7055, train_loss=6.9397]\u001b[A\n", - "Training: 24%|██▎ | 355/1500 [07:44<15:46, 1.21it/s, iou=58.39%, test_loss=8.7055, train_loss=7.1436]\u001b[A\n", - "Training: 24%|██▎ | 356/1500 [07:44<12:38, 1.51it/s, iou=58.39%, test_loss=8.7055, train_loss=7.1436]\u001b[A" - ] - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "16/07/2026-08:41:21.337 INFO:weightslab.components.global_monitoring:pause: \n", - "Training paused.\n", - "16/07/2026-08:41:21.339 INFO:weightslab.trainer.services.experiment_service:ExperimentCommand: [WeightsLab] UI Command: HP changed, experiment paused!\n", - "16/07/2026-08:41:21.340 INFO:weightslab.trainer.services.experiment_service:ExperimentCommand: \n", - "[WeightsLab] UI Command: PAUSE\n", - "16/07/2026-08:41:21.342 INFO:weightslab.components.global_monitoring:pause: \n", - "Training paused.\n", - "16/07/2026-08:41:21.615 INFO:weightslab.components.checkpoint_manager:save_model_checkpoint: Saved model checkpoint: 802319f7f5ee131bc3472a78_step_000055.pt\n", - "16/07/2026-08:41:21.811 INFO:weightslab.components.checkpoint_manager:save_logger_snapshot: Saved logger snapshot: /tmp/tmpf35sd4_z/checkpoints/loggers/loggers.manifest.json (1 chunks)\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "\n", - "Training: 24%|██▎ | 356/1500 [07:45<12:38, 1.51it/s, iou=58.39%, test_loss=8.7055, train_loss=6.4582]\u001b[A\n", - "Training: 24%|██▍ | 357/1500 [07:45<14:34, 1.31it/s, iou=58.39%, test_loss=8.7055, train_loss=6.4582]\u001b[A" - ] - } - ], + "outputs": [], "source": [ "import itertools\n", "\n", diff --git a/weightslab/examples/Notebooks/PyTorch/wl-fraud-detection.ipynb b/weightslab/examples/Notebooks/PyTorch/wl-fraud-detection.ipynb index 920f0bd3..9d104741 100644 --- a/weightslab/examples/Notebooks/PyTorch/wl-fraud-detection.ipynb +++ b/weightslab/examples/Notebooks/PyTorch/wl-fraud-detection.ipynb @@ -174,10 +174,6 @@ "from tqdm.auto import tqdm\n", "\n", "import weightslab as wl\n", - "from weightslab.components.global_monitoring import (\n", - " guard_training_context,\n", - " guard_testing_context,\n", - ")\n", "\n", "logging.basicConfig(level=logging.ERROR)\n", "device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n", @@ -424,7 +420,7 @@ "outputs": [], "source": [ "def train(loader, model, optimizer, criterion, device):\n", - " with guard_training_context:\n", + " with wl.guard_training_context:\n", " inputs, ids, labels = next(loader)\n", " inputs, labels = inputs.to(device), labels.to(device)\n", " optimizer.zero_grad()\n", @@ -440,7 +436,7 @@ "def test(loader, model, criterion, metrics, device, n_batches):\n", " losses = torch.tensor(0.0, device=device)\n", " for inputs, ids, labels in loader:\n", - " with guard_testing_context:\n", + " with wl.guard_testing_context:\n", " inputs, labels = inputs.to(device), labels.to(device)\n", " logits = model(inputs)\n", " preds = logits.argmax(dim=1, keepdim=True)\n", diff --git a/weightslab/examples/Notebooks/PyTorch/wl-segmentation.ipynb b/weightslab/examples/Notebooks/PyTorch/wl-segmentation.ipynb index b1aad59c..00b11eb7 100644 --- a/weightslab/examples/Notebooks/PyTorch/wl-segmentation.ipynb +++ b/weightslab/examples/Notebooks/PyTorch/wl-segmentation.ipynb @@ -201,10 +201,6 @@ "from PIL import Image\n", "\n", "import weightslab as wl\n", - "from weightslab.components.global_monitoring import (\n", - " guard_training_context,\n", - " guard_testing_context,\n", - ")\n", "\n", "# Setup loggers (from main.py)\n", "logging.basicConfig(level=logging.ERROR)\n", @@ -839,7 +835,7 @@ "def train(loader, model, optimizer, sig, metric, device):\n", " \"\"\"Single training step. loader yields (inputs, ids, labels, metadata);\n", " `labels` is a [B, H, W] semantic mask (see seg_collate).\"\"\"\n", - " with guard_training_context:\n", + " with wl.guard_training_context:\n", " (inputs, ids, labels, _) = next(loader)\n", " inputs = inputs.to(device)\n", " labels = labels.to(device) # [B, H, W]\n", @@ -871,7 +867,7 @@ " losses = torch.tensor([0.]).to(device)\n", " dices = torch.tensor([0.]).to(device)\n", " metric.reset()\n", - " with guard_testing_context, torch.no_grad():\n", + " with wl.guard_testing_context, torch.no_grad():\n", " for inputs, ids, labels, _ in loader:\n", " inputs = inputs.to(device)\n", " labels = labels.to(device) # [B, H, W]\n", diff --git a/weightslab/examples/Notebooks/PyTorch/ws-classification.ipynb b/weightslab/examples/Notebooks/PyTorch/ws-classification.ipynb index 9d2286fe..073bbc78 100644 --- a/weightslab/examples/Notebooks/PyTorch/ws-classification.ipynb +++ b/weightslab/examples/Notebooks/PyTorch/ws-classification.ipynb @@ -94,10 +94,6 @@ "\n", "import weightslab as wl\n", "from weightslab.examples.utils.baseline_models.pytorch.models import FashionCNN as CNN\n", - "from weightslab.components.global_monitoring import (\n", - " guard_training_context,\n", - " guard_testing_context,\n", - ")\n", "\n", "logging.basicConfig(level=logging.ERROR)\n", "device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n", @@ -269,7 +265,7 @@ "outputs": [], "source": [ "def train(loader, model, optimizer, criterion, device):\n", - " with guard_training_context:\n", + " with wl.guard_training_context:\n", " inputs, ids, labels = next(loader)\n", " inputs, labels = inputs.to(device), labels.to(device)\n", "\n", @@ -287,7 +283,7 @@ "def test(loader, model, criterion, metric, device, n_batches):\n", " losses = torch.tensor(0.0, device=device)\n", " for inputs, ids, labels in loader:\n", - " with guard_testing_context:\n", + " with wl.guard_testing_context:\n", " inputs, labels = inputs.to(device), labels.to(device)\n", " logits = model(inputs)\n", " preds = logits.argmax(dim=1, keepdim=True)\n", diff --git a/weightslab/examples/Notebooks/Ultralytics/wl-how-to-train-ultralytics-yolo-on-brain-tumor-detection-dataset.ipynb b/weightslab/examples/Notebooks/Ultralytics/wl-how-to-train-ultralytics-yolo-on-brain-tumor-detection-dataset.ipynb index fd85f53c..6e758783 100644 --- a/weightslab/examples/Notebooks/Ultralytics/wl-how-to-train-ultralytics-yolo-on-brain-tumor-detection-dataset.ipynb +++ b/weightslab/examples/Notebooks/Ultralytics/wl-how-to-train-ultralytics-yolo-on-brain-tumor-detection-dataset.ipynb @@ -48,7 +48,7 @@ }, { "cell_type": "code", - "execution_count": 6, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/" @@ -56,24 +56,14 @@ "id": "71be394c", "outputId": "c9ef4767-84f5-4bea-e395-9e11766002f2" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Requirement already satisfied: setuptools in /usr/local/lib/python3.12/dist-packages (75.2.0)\n", - "Requirement already satisfied: wheel in /usr/local/lib/python3.12/dist-packages (0.47.0)\n", - "Requirement already satisfied: packaging>=24.0 in /usr/local/lib/python3.12/dist-packages (from wheel) (26.2)\n" - ] - } - ], + "outputs": [], "source": [ "%pip install setuptools wheel" ] }, { "cell_type": "code", - "execution_count": 7, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/", @@ -82,149 +72,7 @@ "id": "jV1oi_PQN0xK", "outputId": "eda8eca1-f4b3-4ffd-eaa6-0f67ce857153" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Collecting git+https://github.com/GrayboxTech/weightslab.git@landingcollab\n", - " Cloning https://github.com/GrayboxTech/weightslab.git (to revision landingcollab) to /tmp/pip-req-build-z4ba7cx9\n", - " Running command git clone --filter=blob:none --quiet https://github.com/GrayboxTech/weightslab.git /tmp/pip-req-build-z4ba7cx9\n", - " Running command git checkout -b landingcollab --track origin/landingcollab\n", - " Switched to a new branch 'landingcollab'\n", - " Branch 'landingcollab' set up to track remote branch 'landingcollab' from 'origin'.\n", - " Resolved https://github.com/GrayboxTech/weightslab.git to commit 3f1cf143bc93e67df43adc71ab831fa71477e0d2\n", - " Installing build dependencies ... \u001b[?25l\u001b[?25hdone\n", - " Getting requirements to build wheel ... \u001b[?25l\u001b[?25hdone\n", - " Preparing metadata (pyproject.toml) ... \u001b[?25l\u001b[?25hdone\n", - "Requirement already satisfied: numpy<3,>=1.24 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (2.0.2)\n", - "Requirement already satisfied: pandas<3,>=2.2.2 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (2.2.2)\n", - "Requirement already satisfied: duckdb<2,>=1.1 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.3.2)\n", - "Requirement already satisfied: PyYAML<7,>=6.0.3 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (6.0.3)\n", - "Requirement already satisfied: dill<0.5,>=0.3.8 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (0.3.8)\n", - "Requirement already satisfied: zstandard<1,>=0.22 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (0.25.0)\n", - "Requirement already satisfied: h5py<4,>=3.10 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (3.16.0)\n", - "Requirement already satisfied: xxhash<4.1,>=3.4 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (3.7.1)\n", - "Requirement already satisfied: tables<4,>=3.9 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (3.10.2)\n", - "Requirement already satisfied: torch<=2.9,>=2.1 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (2.9.0)\n", - "Requirement already satisfied: torchvision<1,>=0.16 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (0.24.0)\n", - "Requirement already satisfied: torchmetrics>=1.9 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.9.0)\n", - "Requirement already satisfied: grpcio<2,>=1.80 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.81.1)\n", - "Requirement already satisfied: protobuf<8,>=5.28.1 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (5.29.6)\n", - "Requirement already satisfied: pydantic<3,>=2.7 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (2.13.4)\n", - "Requirement already satisfied: Pillow<12,>=10 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (11.3.0)\n", - "Requirement already satisfied: graphviz<1,>=0.20 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (0.21)\n", - "Requirement already satisfied: onnx<=1.20,>=1.15 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.20.0)\n", - "Requirement already satisfied: tqdm<5,>=4.66 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (4.67.3)\n", - "Requirement already satisfied: python-dotenv<2,>=1 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.2.2)\n", - "Requirement already satisfied: langchain-core<2,>=0.3 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.4.9)\n", - "Requirement already satisfied: langchain-ollama<2,>=0.2 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.1.0)\n", - "Requirement already satisfied: langchain-openai<2,>=0.2 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.3.5)\n", - "Requirement already satisfied: typing-extensions~=4.12 in /usr/local/lib/python3.12/dist-packages (from grpcio<2,>=1.80->weightslab==1.3.3.dev8) (4.15.0)\n", - "Requirement already satisfied: jsonpatch<2.0.0,>=1.33.0 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (1.33)\n", - "Requirement already satisfied: langchain-protocol>=0.0.17 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (0.0.18)\n", - "Requirement already satisfied: langsmith<1.0.0,>=0.3.45 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (0.9.1)\n", - "Requirement already satisfied: packaging>=23.2.0 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (26.2)\n", - "Requirement already satisfied: tenacity!=8.4.0,<10.0.0,>=8.1.0 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (9.1.4)\n", - "Requirement already satisfied: uuid-utils<1.0,>=0.12.0 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (0.16.2)\n", - "Requirement already satisfied: ollama<1.0.0,>=0.6.1 in /usr/local/lib/python3.12/dist-packages (from langchain-ollama<2,>=0.2->weightslab==1.3.3.dev8) (0.6.2)\n", - "Requirement already satisfied: openai<3.0.0,>=2.45.0 in /usr/local/lib/python3.12/dist-packages (from langchain-openai<2,>=0.2->weightslab==1.3.3.dev8) (2.45.0)\n", - "Requirement already satisfied: tiktoken<1.0.0,>=0.7.0 in /usr/local/lib/python3.12/dist-packages (from langchain-openai<2,>=0.2->weightslab==1.3.3.dev8) (0.13.0)\n", - "Requirement already satisfied: ml_dtypes>=0.5.0 in /usr/local/lib/python3.12/dist-packages (from onnx<=1.20,>=1.15->weightslab==1.3.3.dev8) (0.5.4)\n", - "Requirement already satisfied: python-dateutil>=2.8.2 in /usr/local/lib/python3.12/dist-packages (from pandas<3,>=2.2.2->weightslab==1.3.3.dev8) (2.9.0.post0)\n", - "Requirement already satisfied: pytz>=2020.1 in /usr/local/lib/python3.12/dist-packages (from pandas<3,>=2.2.2->weightslab==1.3.3.dev8) (2025.2)\n", - "Requirement already satisfied: tzdata>=2022.7 in /usr/local/lib/python3.12/dist-packages (from pandas<3,>=2.2.2->weightslab==1.3.3.dev8) (2026.2)\n", - "Requirement already satisfied: annotated-types>=0.6.0 in /usr/local/lib/python3.12/dist-packages (from pydantic<3,>=2.7->weightslab==1.3.3.dev8) (0.7.0)\n", - "Requirement already satisfied: pydantic-core==2.46.4 in /usr/local/lib/python3.12/dist-packages (from pydantic<3,>=2.7->weightslab==1.3.3.dev8) (2.46.4)\n", - "Requirement already satisfied: typing-inspection>=0.4.2 in /usr/local/lib/python3.12/dist-packages (from pydantic<3,>=2.7->weightslab==1.3.3.dev8) (0.4.2)\n", - "Requirement already satisfied: numexpr>=2.6.2 in /usr/local/lib/python3.12/dist-packages (from tables<4,>=3.9->weightslab==1.3.3.dev8) (2.14.1)\n", - "Requirement already satisfied: py-cpuinfo in /usr/local/lib/python3.12/dist-packages (from tables<4,>=3.9->weightslab==1.3.3.dev8) (9.0.0)\n", - "Requirement already satisfied: blosc2>=2.3.0 in /usr/local/lib/python3.12/dist-packages (from tables<4,>=3.9->weightslab==1.3.3.dev8) (4.5.1)\n", - "Requirement already satisfied: filelock in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.29.4)\n", - "Requirement already satisfied: setuptools in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (75.2.0)\n", - "Requirement already satisfied: sympy>=1.13.3 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (1.14.0)\n", - "Requirement already satisfied: networkx>=2.5.1 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.6.1)\n", - "Requirement already satisfied: jinja2 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.1.6)\n", - "Requirement already satisfied: fsspec>=0.8.5 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (2025.3.0)\n", - "Requirement already satisfied: nvidia-cuda-nvrtc-cu12==12.8.93 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.93)\n", - "Requirement already satisfied: nvidia-cuda-runtime-cu12==12.8.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.90)\n", - "Requirement already satisfied: nvidia-cuda-cupti-cu12==12.8.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.90)\n", - "Requirement already satisfied: nvidia-cudnn-cu12==9.10.2.21 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (9.10.2.21)\n", - "Requirement already satisfied: nvidia-cublas-cu12==12.8.4.1 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.4.1)\n", - "Requirement already satisfied: nvidia-cufft-cu12==11.3.3.83 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (11.3.3.83)\n", - "Requirement already satisfied: nvidia-curand-cu12==10.3.9.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (10.3.9.90)\n", - "Requirement already satisfied: nvidia-cusolver-cu12==11.7.3.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (11.7.3.90)\n", - "Requirement already satisfied: nvidia-cusparse-cu12==12.5.8.93 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.5.8.93)\n", - "Requirement already satisfied: nvidia-cusparselt-cu12==0.7.1 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (0.7.1)\n", - "Requirement already satisfied: nvidia-nccl-cu12==2.27.5 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (2.27.5)\n", - "Requirement already satisfied: nvidia-nvshmem-cu12==3.3.20 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.3.20)\n", - "Requirement already satisfied: nvidia-nvtx-cu12==12.8.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.90)\n", - "Requirement already satisfied: nvidia-nvjitlink-cu12==12.8.93 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.93)\n", - "Requirement already satisfied: nvidia-cufile-cu12==1.13.1.3 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (1.13.1.3)\n", - "Requirement already satisfied: triton==3.5.0 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.5.0)\n", - "Requirement already satisfied: lightning-utilities>=0.15.3 in /usr/local/lib/python3.12/dist-packages (from torchmetrics>=1.9->weightslab==1.3.3.dev8) (0.15.3)\n", - "Requirement already satisfied: ndindex in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (1.10.1)\n", - "Requirement already satisfied: msgpack in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (1.2.1)\n", - "Requirement already satisfied: requests in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (2.32.4)\n", - "Requirement already satisfied: rich in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (13.9.4)\n", - "Requirement already satisfied: textual in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (6.2.1)\n", - "Requirement already satisfied: textual-plotext in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (1.0.1)\n", - "Requirement already satisfied: threadpoolctl in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (3.6.0)\n", - "Requirement already satisfied: jsonpointer>=1.9 in /usr/local/lib/python3.12/dist-packages (from jsonpatch<2.0.0,>=1.33.0->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (3.1.1)\n", - "Requirement already satisfied: anyio>=3.5.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (4.14.0)\n", - "Requirement already satisfied: distro>=1.7.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (1.9.0)\n", - "Requirement already satisfied: httpx<1,>=0.23.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (0.28.1)\n", - "Requirement already satisfied: orjson>=3.9.14 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (3.11.9)\n", - "Requirement already satisfied: requests-toolbelt>=1.0.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (1.0.0)\n", - "Requirement already satisfied: sniffio>=1.1 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (1.3.1)\n", - "Requirement already satisfied: websockets>=15.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (15.0.1)\n", - "Requirement already satisfied: jiter<1,>=0.10.0 in /usr/local/lib/python3.12/dist-packages (from openai<3.0.0,>=2.45.0->langchain-openai<2,>=0.2->weightslab==1.3.3.dev8) (0.15.0)\n", - "Requirement already satisfied: six>=1.5 in /usr/local/lib/python3.12/dist-packages (from python-dateutil>=2.8.2->pandas<3,>=2.2.2->weightslab==1.3.3.dev8) (1.17.0)\n", - "Requirement already satisfied: mpmath<1.4,>=1.1.0 in /usr/local/lib/python3.12/dist-packages (from sympy>=1.13.3->torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (1.3.0)\n", - "Requirement already satisfied: regex in /usr/local/lib/python3.12/dist-packages (from tiktoken<1.0.0,>=0.7.0->langchain-openai<2,>=0.2->weightslab==1.3.3.dev8) (2025.11.3)\n", - "Requirement already satisfied: MarkupSafe>=2.0 in /usr/local/lib/python3.12/dist-packages (from jinja2->torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.0.3)\n", - "Requirement already satisfied: idna>=2.8 in /usr/local/lib/python3.12/dist-packages (from anyio>=3.5.0->langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (3.18)\n", - "Requirement already satisfied: certifi in /usr/local/lib/python3.12/dist-packages (from httpx<1,>=0.23.0->langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (2026.6.17)\n", - "Requirement already satisfied: httpcore==1.* in /usr/local/lib/python3.12/dist-packages (from httpx<1,>=0.23.0->langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (1.0.9)\n", - "Requirement already satisfied: h11>=0.16 in /usr/local/lib/python3.12/dist-packages (from httpcore==1.*->httpx<1,>=0.23.0->langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (0.16.0)\n", - "Requirement already satisfied: charset_normalizer<4,>=2 in /usr/local/lib/python3.12/dist-packages (from requests->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (3.4.7)\n", - "Requirement already satisfied: urllib3<3,>=1.21.1 in /usr/local/lib/python3.12/dist-packages (from requests->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (2.5.0)\n", - "Requirement already satisfied: markdown-it-py>=2.2.0 in /usr/local/lib/python3.12/dist-packages (from rich->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (4.2.0)\n", - "Requirement already satisfied: pygments<3.0.0,>=2.13.0 in /usr/local/lib/python3.12/dist-packages (from rich->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (2.20.0)\n", - "Requirement already satisfied: platformdirs<5,>=3.6.0 in /usr/local/lib/python3.12/dist-packages (from textual->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (4.10.0)\n", - "Requirement already satisfied: plotext<6.0.0,>=5.2.8 in /usr/local/lib/python3.12/dist-packages (from textual-plotext->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (5.3.2)\n", - "Requirement already satisfied: mdurl~=0.1 in /usr/local/lib/python3.12/dist-packages (from markdown-it-py>=2.2.0->rich->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (0.1.2)\n", - "Requirement already satisfied: linkify-it-py<3,>=1 in /usr/local/lib/python3.12/dist-packages (from markdown-it-py[linkify,plugins]>=2.1.0->textual->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (2.1.0)\n", - "Requirement already satisfied: mdit-py-plugins>=0.5.0 in /usr/local/lib/python3.12/dist-packages (from markdown-it-py[linkify,plugins]>=2.1.0->textual->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (0.6.1)\n", - "Requirement already satisfied: uc-micro-py in /usr/local/lib/python3.12/dist-packages (from linkify-it-py<3,>=1->markdown-it-py[linkify,plugins]>=2.1.0->textual->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (2.0.0)\n", - "Building wheels for collected packages: weightslab\n", - " Building wheel for weightslab (pyproject.toml) ... \u001b[?25l\u001b[?25hdone\n", - " Created wheel for weightslab: filename=weightslab-1.3.3.dev8-py3-none-any.whl size=2827715 sha256=5c4e7c14bbf95a48b9a8c86b5a8c744af67e1eb395197cbe79d3e443c21cf86e\n", - " Stored in directory: /tmp/pip-ephem-wheel-cache-__zwgf7x/wheels/4d/c7/cd/817c2745790555e7b6c532f5b6055d47567f9416fdfc65879b\n", - "Successfully built weightslab\n", - "Installing collected packages: weightslab\n", - " Attempting uninstall: weightslab\n", - " Found existing installation: weightslab 1.3.3.dev7\n", - " Uninstalling weightslab-1.3.3.dev7:\n", - " Successfully uninstalled weightslab-1.3.3.dev7\n", - "Successfully installed weightslab-1.3.3.dev8\n" - ] - }, - { - "data": { - "application/vnd.colab-display-data+json": { - "id": "b7fce274042740d19087972fe9c01c9b", - "pip_warning": { - "packages": [ - "weightslab" - ] - } - } - }, - "metadata": {}, - "output_type": "display_data" - } - ], + "outputs": [], "source": [ "# Install WeightsLab. Colab already ships torch, torchvision, numpy, scikit-learn\n", "# and Pillow, so nothing extra is needed here.\n", @@ -234,7 +82,7 @@ }, { "cell_type": "code", - "execution_count": 1, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/" @@ -242,39 +90,7 @@ "id": "pJhvREbaKs7f", "outputId": "651baa36-9ebb-49da-830b-e6c8a6e2b64f" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Ultralytics 8.4.96 🚀 Python-3.12.13 torch-2.9.0+cu128 CPU (Intel Xeon CPU @ 2.20GHz)\n", - "Setup complete ✅ (2 CPUs, 12.7 GB RAM, 30.5/107.7 GB disk)\n", - "16/07/2026-09:29:29.683 INFO:root:setup_logging: WeightsLab logging initialized - Log file: /tmp/tmp6meuzqdz/weightslab_logs/weightslab_20260716_092929.log\n", - "16/07/2026-09:29:29.685 INFO:weightslab:: WeightsLab package initialized - Log level: INFO, Log to file: True\n", - "16/07/2026-09:29:29.686 INFO:weightslab:: \n", - "╭─────────────────────────────────────────────────────────────────────────────────────────────────────╮\n", - "│ │\n", - "│ \u001b[32m$$\\ $$\\ \u001b[0m $$\\ $$\\ $$\\ \u001b[31m$$\\ \u001b[0m $$\\ │\n", - "│ \u001b[32m$$ | $\\ $$ |\u001b[0m \\__| $$ | $$ | \u001b[31m$$ | \u001b[0m $$ | │\n", - "│ \u001b[32m$$ |$$$\\ $$ |\u001b[0m $$$$$$\\ $$\\ $$$$$$\\ $$$$$$$\\ $$$$$$\\ $$$$$$$\\ \u001b[31m$$ | \u001b[0m $$$$$$\\ $$$$$$$\\ │\n", - "│ \u001b[32m$$ $$ $$\\$$ |\u001b[0m$$ __$$\\ $$ |$$ __$$\\ $$ __$$\\ \\_$$ _| $$ _____|\u001b[31m$$ | \u001b[0m \\____$$\\ $$ __$$\\ │\n", - "│ \u001b[32m$$$$ _$$$$ |\u001b[0m$$$$$$$$ |$$ |$$ / $$ |$$ | $$ | $$ | \\$$$$$$\\ \u001b[31m$$ | \u001b[0m $$$$$$$ |$$ | $$ | │\n", - "│ \u001b[32m$$$ / \\$$$ |\u001b[0m$$ ____|$$ |$$ | $$ |$$ | $$ | $$ |$$\\ \\____$$\\ \u001b[31m$$ | \u001b[0m$$ __$$ |$$ | $$ | │\n", - "│ \u001b[32m$$ / \\$$ |\u001b[0m\\$$$$$$$\\ $$ |\\$$$$$$$ |$$ | $$ | \\$$$$ |$$$$$$$ |\u001b[31m$$$$$$$$\\ \u001b[0m\\$$$$$$$ |$$$$$$$ | │\n", - "│ \u001b[32m\\__/ \\__|\u001b[0m \\_______|\\__| \\____$$ |\\__| \\__| \\____/ \\_______/ \u001b[31m\\________|\u001b[0m \\_______|\\_______/ │\n", - "│ \u001b[32m \u001b[0m $$\\ $$ |\u001b[31m\u001b[0m │\n", - "│ \u001b[32m \u001b[0m \\$$$$$$ |\u001b[31m\u001b[0m │\n", - "│ \u001b[32m \u001b[0m \\______/\u001b[31m\u001b[0m │\n", - "│ │\n", - "│ Inspect - Edit - Evolve Neural Networks │\n", - "│ │\n", - "╰─── By GrayBx, v1.3.3.dev8 ──────────────────────────────────────────────────────────────────────────╯\n", - "\n", - "\n", - "16/07/2026-09:29:29.690 INFO:weightslab.utils.telemetry:ping_import: WeightsLab uses anonymous usage data (package version used, and OS name). Set WL_NO_TELEMETRY=1 to disable.\n" - ] - } - ], + "outputs": [], "source": [ "import ultralytics\n", "ultralytics.checks()\n", @@ -321,7 +137,7 @@ }, { "cell_type": "code", - "execution_count": 2, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/", @@ -330,218 +146,7 @@ "id": "Svz7NK9qKs7k", "outputId": "2931e6b9-a06e-4dc0-9f67-e2e6216ee1c9" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "16/07/2026-09:29:52.825 INFO:weightslab.src:watch_or_edit: LoggerQueue for experiment history has been initialized and registered.\n", - "16/07/2026-09:29:52.828 INFO:weightslab.components.checkpoint_manager:__init__: CheckpointManager initialized at /tmp/tmp3abnx3ay\n", - "16/07/2026-09:29:52.829 INFO:weightslab.src:watch_or_edit: Registered new checkpoint manager in ledger\n", - "16/07/2026-09:29:52.830 INFO:root:set_log_directory: Log file moved from /tmp/tmp6meuzqdz/weightslab_logs/weightslab_20260716_092929.log to /tmp/tmp3abnx3ay/weightslab_20260716_092929.log\n", - "16/07/2026-09:29:52.832 INFO:root:set_log_directory: Log directory updated to: /tmp/tmp3abnx3ay\n", - "16/07/2026-09:29:52.833 INFO:root:set_log_directory: Log file: /tmp/tmp3abnx3ay/weightslab_20260716_092929.log\n", - "16/07/2026-09:29:52.835 INFO:weightslab.trainer.trainer_services:_run_security_preflight: \n", - "\n", - "# #######################################\n", - "# #######################################\n", - "[gRPC] Security preflight checks:\n", - "\tTLS: DISABLED (unencrypted traffic)\n", - "\tAuth tokens: NONE configured\n", - "\t! WARNING: GRPC_TLS_ENABLED=0. Traffic will be unencrypted. Use only for development.\n", - "\t! WARNING: No GRPC_AUTH_TOKEN/GRPC_AUTH_TOKENS configured. Only transport-level trust (TLS/mTLS) will protect RPC access.\n", - "# #######################################\n", - "# #######################################\n", - "\n", - "16/07/2026-09:29:52.837 INFO:weightslab.trainer.trainer_services:grpc_serve: [gRPC] Watchdogs disabled via WEIGHTSLAB_DISABLE_WATCHDOGS.\n", - "16/07/2026-09:29:52.839 INFO:weightslab.trainer.trainer_services:serving_thread_callback: [gRPC] Thread callback started\n", - "16/07/2026-09:29:52.840 INFO:weightslab.trainer.trainer_services:grpc_serve: [gRPC] Server started with watchdogs disabled (host=0.0.0.0 port=50051 workers=None)\n", - "16/07/2026-09:29:52.841 INFO:weightslab.trainer.trainer_services:serving_thread_callback: [gRPC] Creating ThreadPoolExecutor with 6 worker threads (n_workers_grpc=None, max_concurrent_rpcs=None)\n", - "16/07/2026-09:29:52.842 WARNING:weightslab.src:serve: Running in a notebook/Colab with serving_grpc=True but serving_bore=False. Weights Studio runs on your OWN machine and cannot reach this backend directly. Pass serving_bore=True to open a tunnel and sync with the UI: wl.serve(serving_grpc=True, serving_bore=True).\n", - "16/07/2026-09:29:52.852 INFO:weightslab.trainer.trainer_services:serving_thread_callback: [gRPC] Server object created\n", - "16/07/2026-09:29:52.854 INFO:weightslab.backend.cli:cli_serve: cli_thread_started\n", - "16/07/2026-09:29:52.859 INFO:weightslab.trainer.services.agent.agent:__init__: Initializing DataManipulationAgent\n", - "16/07/2026-09:29:52.860 INFO:weightslab.trainer.services.agent.agent:_setup_model_schema: [Agent] components.get('model') returned None (ledger-registered model names: []). Model-related agent requests will report 'no model registered'. If a model IS registered under a name other than the experiment name / 'experiment' / 'main', ExperimentContext.ensure_components()'s model-resolution heuristic won't find it.\n", - "16/07/2026-09:29:52.867 WARNING:weightslab.trainer.services.agent.agent:_load_config: Error loading config from /usr/local/lib/python3.12/dist-packages: [Errno 21] Is a directory: '/usr/local/lib/python3.12/dist-packages'\n", - "16/07/2026-09:29:52.869 INFO:weightslab.trainer.services.agent.agent:_load_config: \n", - "\n", - "# #######################################\n", - "# #######################################\n", - "Agent initialized from configuration /content/agent_config.yaml: \n", - "\tFinal Agent Configuration: Preferred Provider=openrouter, \n", - "\tFallback to Local=True, \n", - "\tOpenRouter Model=~google/gemini-flash-latest with:\n", - "\t\tAPI Key=None\n", - "\t\tBase URL=https://openrouter.ai/api/v1, \n", - "\tOllama Model=llama3.2:3b\n", - "# #######################################\n", - "# #######################################\n", - "\n", - "16/07/2026-09:29:52.870 INFO:weightslab.trainer.services.agent.agent:_setup_providers: Setting up Ollama with model llama3.2:3b\n", - "16/07/2026-09:29:52.915 INFO:weightslab.components.global_monitoring:resume: \n", - "Attempting to resume training...\n", - "16/07/2026-09:29:52.916 INFO:weightslab.components.checkpoint_manager:update_experiment_hash: First time initialization; skipping hash update.\n", - "16/07/2026-09:29:52.917 INFO:weightslab.components.experiment_hash:has_changed: Experiment configuration changed: {'hp'}\n", - "16/07/2026-09:29:52.918 INFO:weightslab.components.experiment_hash:generate_hash: Generated experiment hash: 1b111ac40000000000000000- (HP: 1b111ac4, Model: 00000000, Data: 00000000)\n", - "16/07/2026-09:29:52.919 INFO:weightslab.components.checkpoint_manager:update_experiment_hash: Initial experiment hash set: 1b111ac4-00000000-00000000\n", - "16/07/2026-09:29:52.919 INFO:weightslab.components.checkpoint_manager:update_experiment_hash: Changed components: {'hp'}\n", - "16/07/2026-09:29:52.922 INFO:weightslab.components.checkpoint_manager:_save_changes: Dumping hyperparameters config...\n", - "16/07/2026-09:29:52.924 INFO:weightslab.components.checkpoint_manager:save_config: Saved config: 1b111ac4_config.yaml with exp_hash prefix: 1b111ac4\n", - "16/07/2026-09:29:52.925 WARNING:weightslab.components.checkpoint_manager:_save_changes: Could not save weights: no model available\n", - "16/07/2026-09:29:52.931 INFO:weightslab.components.checkpoint_manager:save_logger_snapshot: Saved logger snapshot: /tmp/tmp3abnx3ay/checkpoints/loggers/loggers.manifest.json (1 chunks)\n", - "16/07/2026-09:29:52.935 INFO:weightslab.components.global_monitoring:resume: Hashes by module: ['1b111ac4', '00000000', '00000000']\n", - "16/07/2026-09:29:52.935 WARNING:weightslab.components.global_monitoring:resume: Cannot resume training: experiment hash not computed yet for every modules ['1b111ac4', '00000000', '00000000'].\n", - "16/07/2026-09:29:53.028 INFO:weightslab.trainer.services.agent.agent:_setup_providers: [Agent] Ollama enabled: llama3.2:3b\n", - "16/07/2026-09:29:53.031 INFO:weightslab.trainer.services.data_service:__init__: DataService initialized.\n", - "16/07/2026-09:29:53.036 INFO:weightslab.trainer.trainer_services:serving_thread_callback: [gRPC] Servicer added\n", - "16/07/2026-09:29:53.040 INFO:weightslab.trainer.trainer_services:serving_thread_callback: [gRPC] Attempting to bind to 0.0.0.0:50051\n", - "16/07/2026-09:29:53.044 WARNING:weightslab.trainer.trainer_services:serving_thread_callback: [gRPC] TLS disabled; using insecure transport on 0.0.0.0:50051\n", - "16/07/2026-09:29:53.046 INFO:weightslab.trainer.trainer_services:serving_thread_callback: [gRPC] Port 50051 bound successfully.\n", - "16/07/2026-09:29:53.051 INFO:weightslab.trainer.trainer_services:serving_thread_callback: [gRPC] Server started and listening on 0.0.0.0:50051\n", - "Ultralytics 8.4.96 🚀 Python-3.12.13 torch-2.9.0+cu128 CPU (Intel Xeon CPU @ 2.20GHz)\n", - "\u001b[34m\u001b[1mengine/trainer: \u001b[0magnostic_nms=False, amp=False, angle=1.0, augment=False, auto_augment=None, batch=16, bgr=0.0, box=7.5, cache=False, cfg=None, classes=None, close_mosaic=10, cls=0.5, cls_pw=0.0, cls_remap=True, compile=False, conf=None, copy_paste=0.0, copy_paste_mode=flip, cos_lr=False, cutmix=0.0, data=/content/datasets/brain-tumor/brain-tumor.yaml, degrees=0.0, deterministic=True, device=cpu, dfl=1.5, dis=6.0, distill_model=None, dnn=False, dropout=0.0, dynamic=False, embed=None, end2end=None, epochs=10, erasing=0.0, exist_ok=False, fliplr=0.0, flipud=0.0, format=torchscript, fraction=1.0, freeze=None, hsv_h=0.0, hsv_s=0.0, hsv_v=0.0, imgsz=640, iou=0.7, keras=False, kobj=1.0, line_width=None, lr0=0.001, lrf=0.01, mask_ratio=4, max_det=300, mixup=0.0, mode=train, model=yolo11n.pt, momentum=0.937, mosaic=0.0, multi_scale=0.0, name=brain-tumor-4, nbs=64, nms=False, opset=None, optimize=False, optimizer=SGD, overlap_mask=True, patience=100, perspective=0.0, plots=True, pose=12.0, pretrained=True, profile=False, project=./wl_logs, quantize=None, rect=False, resume=False, retina_masks=False, rle=1.0, save=True, save_conf=False, save_crop=False, save_dir=/content/runs/detect/wl_logs/brain-tumor-4, save_frames=False, save_json=False, save_period=-1, save_txt=False, scale=0.0, seed=0, shear=0.0, show=False, show_boxes=True, show_conf=True, show_labels=True, simplify=True, single_cls=False, source=None, split=val, stream_buffer=False, task=detect, time=None, tracker=tracktrack.yaml, translate=0.0, val=True, verbose=True, vid_stride=1, visualize=False, warmup_bias_lr=0.1, warmup_epochs=3.0, warmup_momentum=0.8, weight_decay=0.0005, workers=0, workspace=None\n", - "Overriding model.yaml nc=80 with nc=2\n", - "\n", - " from n params module arguments \n", - " 0 -1 1 464 ultralytics.nn.modules.conv.Conv [3, 16, 3, 2] \n", - " 1 -1 1 4672 ultralytics.nn.modules.conv.Conv [16, 32, 3, 2] \n", - " 2 -1 1 6640 ultralytics.nn.modules.block.C3k2 [32, 64, 1, False, 0.25] \n", - " 3 -1 1 36992 ultralytics.nn.modules.conv.Conv [64, 64, 3, 2] \n", - " 4 -1 1 26080 ultralytics.nn.modules.block.C3k2 [64, 128, 1, False, 0.25] \n", - " 5 -1 1 147712 ultralytics.nn.modules.conv.Conv [128, 128, 3, 2] \n", - " 6 -1 1 87040 ultralytics.nn.modules.block.C3k2 [128, 128, 1, True] \n", - " 7 -1 1 295424 ultralytics.nn.modules.conv.Conv [128, 256, 3, 2] \n", - " 8 -1 1 346112 ultralytics.nn.modules.block.C3k2 [256, 256, 1, True] \n", - " 9 -1 1 164608 ultralytics.nn.modules.block.SPPF [256, 256, 5] \n", - " 10 -1 1 249728 ultralytics.nn.modules.block.C2PSA [256, 256, 1] \n", - " 11 -1 1 0 torch.nn.modules.upsampling.Upsample [None, 2, 'nearest'] \n", - " 12 [-1, 6] 1 0 ultralytics.nn.modules.conv.Concat [1] \n", - " 13 -1 1 111296 ultralytics.nn.modules.block.C3k2 [384, 128, 1, False] \n", - " 14 -1 1 0 torch.nn.modules.upsampling.Upsample [None, 2, 'nearest'] \n", - " 15 [-1, 4] 1 0 ultralytics.nn.modules.conv.Concat [1] \n", - " 16 -1 1 32096 ultralytics.nn.modules.block.C3k2 [256, 64, 1, False] \n", - " 17 -1 1 36992 ultralytics.nn.modules.conv.Conv [64, 64, 3, 2] \n", - " 18 [-1, 13] 1 0 ultralytics.nn.modules.conv.Concat [1] \n", - " 19 -1 1 86720 ultralytics.nn.modules.block.C3k2 [192, 128, 1, False] \n", - " 20 -1 1 147712 ultralytics.nn.modules.conv.Conv [128, 128, 3, 2] \n", - " 21 [-1, 10] 1 0 ultralytics.nn.modules.conv.Concat [1] \n", - " 22 -1 1 378880 ultralytics.nn.modules.block.C3k2 [384, 256, 1, True] \n", - " 23 [16, 19, 22] 1 431062 ultralytics.nn.modules.head.Detect [2, 16, None, [64, 128, 256]] \n", - "YOLO11n summary: 182 layers, 2,590,230 parameters, 2,590,214 gradients, 6.4 GFLOPs\n", - "\n", - "Transferred 448/499 items from pretrained weights\n", - "Freezing layer 'model.23.dfl.conv.weight'\n", - "\u001b[34m\u001b[1mtrain: \u001b[0mFast image access ✅ (ping: 0.0±0.0 ms, read: 115.3±57.6 MB/s, size: 3.6 KB)\n", - "\u001b[K\u001b[34m\u001b[1mtrain: \u001b[0mScanning /content/datasets/brain-tumor/labels/train.cache... 878 images, 15 backgrounds, 0 corrupt: 100% ━━━━━━━━━━━━ 893/893 101.2Mit/s 0.0s\n", - "\u001b[34m\u001b[1malbumentations: \u001b[0mBlur(p=0.01, blur_limit=(3, 7)), MedianBlur(p=0.01, blur_limit=(3, 7)), ToGray(p=0.01, method='weighted_average', num_output_channels=3), CLAHE(p=0.01, clip_limit=(1.0, 4.0), tile_grid_size=(8, 8))\n", - "16/07/2026-09:29:54.526 INFO:weightslab.data.data_samples_with_ops:__init__: [DataSampleTrackingWrapper] H5 persistence enabled at /tmp/tmp3abnx3ay/checkpoints/data/data.h5\n", - "16/07/2026-09:29:54.529 INFO:weightslab.data.data_samples_with_ops:__init__: Preloading sample statistics for PHYSICAL indices in split 'train_loader' with: preload_labels=True, preload_metadata=True...\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "Initializing ledger for split 'train_loader': 100%|██████████| 893/893 [00:00<00:00, 13970.17it/s]" - ] - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "16/07/2026-09:29:54.600 INFO:weightslab.data.dataframe_manager:register_split: [LedgeredDataFrameManager] Registering split 'train_loader' with 893 samples.\n", - "16/07/2026-09:29:54.615 INFO:weightslab.data.dataframe_manager:register_split: [LedgeredDataFrameManager] After annotation expansion: 893 samples → 984 annotation rows.\n", - "16/07/2026-09:29:54.616 INFO:weightslab.data.h5_array_store:__init__: [H5ArrayStore] Initialized with cache limit: 2048MB\n", - "16/07/2026-09:29:54.636 INFO:weightslab.components.checkpoint_manager:load_checkpoint: Loading checkpoint 1b111ac400000000...\n", - "16/07/2026-09:29:54.637 INFO:weightslab.components.checkpoint_manager:load_checkpoint: Target: HP=1b111ac4 MODEL=00000000 DATA=00000000\n", - "16/07/2026-09:29:54.638 INFO:weightslab.components.checkpoint_manager:load_checkpoint: Current: HP=1b111ac4 MODEL=00000000 DATA=00000000\n", - "16/07/2026-09:29:54.639 INFO:weightslab.components.checkpoint_manager:load_checkpoint: [-] Model architecture unchanged, using current model\n", - "16/07/2026-09:29:54.639 INFO:weightslab.components.checkpoint_manager:load_checkpoint: [-] Config unchanged, using current config\n", - "16/07/2026-09:29:54.641 WARNING:weightslab.components.checkpoint_manager:load_checkpoint: [WARNING] Data snapshot file not found: /tmp/tmp3abnx3ay/checkpoints/data/00000000/00000000_data_snapshot.json\n", - "16/07/2026-09:29:54.641 INFO:weightslab.components.checkpoint_manager:load_checkpoint: Loaded components: set()\n", - "\u001b[34m\u001b[1mval: \u001b[0mFast image access ✅ (ping: 0.0±0.0 ms, read: 62.7±38.0 MB/s, size: 3.9 KB)\n", - "\u001b[K\u001b[34m\u001b[1mval: \u001b[0mScanning /content/datasets/brain-tumor/labels/val.cache... 223 images, 0 backgrounds, 0 corrupt: 100% ━━━━━━━━━━━━ 223/223 32.3Mit/s 0.0s\n", - "16/07/2026-09:29:54.660 INFO:weightslab.data.data_samples_with_ops:__init__: [DataSampleTrackingWrapper] H5 persistence enabled at /tmp/tmp3abnx3ay/checkpoints/data/data.h5\n", - "16/07/2026-09:29:54.662 INFO:weightslab.data.data_samples_with_ops:__init__: Preloading sample statistics for PHYSICAL indices in split 'val_loader' with: preload_labels=True, preload_metadata=True...\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "\n", - "Initializing ledger for split 'val_loader': 100%|██████████| 223/223 [00:00<00:00, 5362.24it/s]" - ] - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "16/07/2026-09:29:54.708 INFO:weightslab.data.dataframe_manager:register_split: [LedgeredDataFrameManager] Registering split 'val_loader' with 223 samples.\n", - "16/07/2026-09:29:54.714 INFO:weightslab.data.dataframe_manager:register_split: [LedgeredDataFrameManager] After annotation expansion: 223 samples → 257 annotation rows.\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "\n" - ] - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "16/07/2026-09:29:54.734 INFO:weightslab.components.checkpoint_manager:load_checkpoint: Loading checkpoint 1b111ac400000000...\n", - "16/07/2026-09:29:54.736 INFO:weightslab.components.checkpoint_manager:load_checkpoint: Target: HP=1b111ac4 MODEL=00000000 DATA=00000000\n", - "16/07/2026-09:29:54.739 INFO:weightslab.components.checkpoint_manager:load_checkpoint: Current: HP=1b111ac4 MODEL=00000000 DATA=00000000\n", - "16/07/2026-09:29:54.740 INFO:weightslab.components.checkpoint_manager:load_checkpoint: [-] Model architecture unchanged, using current model\n", - "16/07/2026-09:29:54.742 INFO:weightslab.components.checkpoint_manager:load_checkpoint: [-] Config unchanged, using current config\n", - "16/07/2026-09:29:54.743 WARNING:weightslab.components.checkpoint_manager:load_checkpoint: [WARNING] Data snapshot file not found: /tmp/tmp3abnx3ay/checkpoints/data/00000000/00000000_data_snapshot.json\n", - "16/07/2026-09:29:54.746 INFO:weightslab.components.checkpoint_manager:load_checkpoint: Loaded components: set()\n", - "\u001b[34m\u001b[1moptimizer:\u001b[0m SGD(lr=0.001, momentum=0.937) with parameter groups 81 weight(decay=0.0), 88 weight(decay=0.0005), 87 bias(decay=0.0)\n", - "Plotting labels to /content/runs/detect/wl_logs/brain-tumor-4/labels.jpg... \n", - "16/07/2026-09:29:55.639 INFO:weightslab.backend.model_interface:__init__: Using checkpoint manager from ledger\n", - "16/07/2026-09:29:55.644 INFO:weightslab.components.checkpoint_manager:load_checkpoint: Loading checkpoint 1b111ac400000000...\n", - "16/07/2026-09:29:55.645 INFO:weightslab.components.checkpoint_manager:load_checkpoint: Target: HP=1b111ac4 MODEL=00000000 DATA=00000000\n", - "16/07/2026-09:29:55.645 INFO:weightslab.components.checkpoint_manager:load_checkpoint: Current: HP=1b111ac4 MODEL=00000000 DATA=00000000\n", - "16/07/2026-09:29:55.646 INFO:weightslab.components.checkpoint_manager:load_checkpoint: [-] Model architecture unchanged, using current model\n", - "16/07/2026-09:29:55.647 WARNING:weightslab.components.checkpoint_manager:load_checkpoint: [WARNING] No weight files found for 00000000\n", - "16/07/2026-09:29:55.648 INFO:weightslab.components.checkpoint_manager:load_checkpoint: [-] Config unchanged, using current config\n", - "16/07/2026-09:29:55.648 INFO:weightslab.components.checkpoint_manager:load_checkpoint: Loaded components: set()\n", - "Image sizes 640 train, 640 val\n", - "Using 0 dataloader workers\n", - "Logging results to \u001b[1m/content/runs/detect/wl_logs/brain-tumor-4\u001b[0m\n", - "Starting training for 10 epochs...\n", - "16/07/2026-09:29:55.661 INFO:weightslab.backend.dataloader_interface:set_batch_size: Batch size updated: 16 -> 4 (Loader: train_loader)\n", - "Closing dataloader mosaic\n", - "\u001b[34m\u001b[1malbumentations: \u001b[0mBlur(p=0.01, blur_limit=(3, 7)), MedianBlur(p=0.01, blur_limit=(3, 7)), ToGray(p=0.01, method='weighted_average', num_output_channels=3), CLAHE(p=0.01, clip_limit=(1.0, 4.0), tile_grid_size=(8, 8))\n", - "\n", - " Epoch GPU_mem box_loss cls_loss dfl_loss Instances Size\n", - "\u001b[K: 0% ──────────── 0/56 1:03\n" - ] - }, - { - "ename": "KeyboardInterrupt", - "evalue": "", - "output_type": "error", - "traceback": [ - "\u001b[0;31m---------------------------------------------------------------------------\u001b[0m", - "\u001b[0;31mKeyboardInterrupt\u001b[0m Traceback (most recent call last)", - "\u001b[0;32m/tmp/ipykernel_6805/1732828701.py\u001b[0m in \u001b[0;36m\u001b[0;34m()\u001b[0m\n\u001b[1;32m 30\u001b[0m \u001b[0;31m# Read back the (now live) hyperparameters and hand training to WLAwareTrainer.\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 31\u001b[0m \u001b[0mmodel\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mYOLO\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mcfg\u001b[0m\u001b[0;34m[\u001b[0m\u001b[0;34m\"model\"\u001b[0m\u001b[0;34m]\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m---> 32\u001b[0;31m results = model.train(\n\u001b[0m\u001b[1;32m 33\u001b[0m \u001b[0mtrainer\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0mWLAwareTrainer\u001b[0m\u001b[0;34m,\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 34\u001b[0m \u001b[0mdata\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0mstr\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mcfg\u001b[0m\u001b[0;34m[\u001b[0m\u001b[0;34m\"data\"\u001b[0m\u001b[0;34m]\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m,\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n", - "\u001b[0;32m/usr/local/lib/python3.12/dist-packages/ultralytics/engine/model.py\u001b[0m in \u001b[0;36mtrain\u001b[0;34m(self, trainer, **kwargs)\u001b[0m\n\u001b[1;32m 814\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mmodel\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mtrainer\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mmodel\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 815\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 816\u001b[0;31m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mtrainer\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mtrain\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m 817\u001b[0m \u001b[0;31m# Update model and cfg after training\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 818\u001b[0m \u001b[0;32mif\u001b[0m \u001b[0mRANK\u001b[0m \u001b[0;32min\u001b[0m \u001b[0;34m{\u001b[0m\u001b[0;34m-\u001b[0m\u001b[0;36m1\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;36m0\u001b[0m\u001b[0;34m}\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n", - "\u001b[0;32m/usr/local/lib/python3.12/dist-packages/ultralytics/engine/trainer.py\u001b[0m in \u001b[0;36mtrain\u001b[0;34m(self)\u001b[0m\n\u001b[1;32m 239\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 240\u001b[0m \u001b[0;32melse\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 241\u001b[0;31m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_do_train\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m 242\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 243\u001b[0m \u001b[0;32mdef\u001b[0m \u001b[0m_setup_scheduler\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mself\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n", - "\u001b[0;32m/usr/local/lib/python3.12/dist-packages/ultralytics/engine/trainer.py\u001b[0m in \u001b[0;36m_do_train\u001b[0;34m(self)\u001b[0m\n\u001b[1;32m 430\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mtloss\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0;32mNone\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 431\u001b[0m \u001b[0;32mfor\u001b[0m \u001b[0mi\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mbatch\u001b[0m \u001b[0;32min\u001b[0m \u001b[0mpbar\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 432\u001b[0;31m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mrun_callbacks\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m\"on_train_batch_start\"\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m 433\u001b[0m \u001b[0;31m# Warmup\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 434\u001b[0m \u001b[0mni\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mi\u001b[0m \u001b[0;34m+\u001b[0m \u001b[0mnb\u001b[0m \u001b[0;34m*\u001b[0m \u001b[0mepoch\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n", - "\u001b[0;32m/usr/local/lib/python3.12/dist-packages/ultralytics/engine/trainer.py\u001b[0m in \u001b[0;36mrun_callbacks\u001b[0;34m(self, event)\u001b[0m\n\u001b[1;32m 210\u001b[0m \u001b[0;34m\"\"\"Run all existing callbacks associated with a particular event.\"\"\"\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 211\u001b[0m \u001b[0;32mfor\u001b[0m \u001b[0mcallback\u001b[0m \u001b[0;32min\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mcallbacks\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mget\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mevent\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m[\u001b[0m\u001b[0;34m]\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 212\u001b[0;31m \u001b[0mcallback\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mself\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m 213\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 214\u001b[0m \u001b[0;32mdef\u001b[0m \u001b[0mtrain\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mself\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n", - "\u001b[0;32m/usr/local/lib/python3.12/dist-packages/weightslab/integrations/ultralytics/trainer.py\u001b[0m in \u001b[0;36m_on_train_batch_start\u001b[0;34m(trainer)\u001b[0m\n\u001b[1;32m 97\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 98\u001b[0m \u001b[0;32mdef\u001b[0m \u001b[0m_on_train_batch_start\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mtrainer\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m---> 99\u001b[0;31m \u001b[0mwl\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mguard_training_context\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m__enter__\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m 100\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 101\u001b[0m \u001b[0;32mdef\u001b[0m \u001b[0m_on_train_batch_end\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mtrainer\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n", - "\u001b[0;32m/usr/local/lib/python3.12/dist-packages/weightslab/components/global_monitoring.py\u001b[0m in \u001b[0;36m__enter__\u001b[0;34m(self, f)\u001b[0m\n\u001b[1;32m 224\u001b[0m \u001b[0;32mif\u001b[0m \u001b[0mf\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 225\u001b[0m \u001b[0mpause_controller\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mresume\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mforce\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0mf\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 226\u001b[0;31m \u001b[0mpause_controller\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mwait_if_paused\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m 227\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0marchitecture_guard\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m__enter__\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 228\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n", - "\u001b[0;32m/usr/local/lib/python3.12/dist-packages/weightslab/components/global_monitoring.py\u001b[0m in \u001b[0;36mwait_if_paused\u001b[0;34m(self, skip_pause)\u001b[0m\n\u001b[1;32m 119\u001b[0m \u001b[0;31m# Also wakes up early when an evaluation is pending/running so the\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 120\u001b[0m \u001b[0;31m# training loop and dataloaders can service evaluation mode.\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 121\u001b[0;31m \u001b[0;32mwhile\u001b[0m \u001b[0;32mnot\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_event\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mwait\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mtimeout\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0;36m0.5\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m 122\u001b[0m \u001b[0;31m# Timeout occurred – check for evaluation request before looping\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 123\u001b[0m \u001b[0;32mtry\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n", - "\u001b[0;32m/usr/lib/python3.12/threading.py\u001b[0m in \u001b[0;36mwait\u001b[0;34m(self, timeout)\u001b[0m\n\u001b[1;32m 653\u001b[0m \u001b[0msignaled\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_flag\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 654\u001b[0m \u001b[0;32mif\u001b[0m \u001b[0;32mnot\u001b[0m \u001b[0msignaled\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 655\u001b[0;31m \u001b[0msignaled\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_cond\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mwait\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mtimeout\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m 656\u001b[0m \u001b[0;32mreturn\u001b[0m \u001b[0msignaled\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 657\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n", - "\u001b[0;32m/usr/lib/python3.12/threading.py\u001b[0m in \u001b[0;36mwait\u001b[0;34m(self, timeout)\u001b[0m\n\u001b[1;32m 357\u001b[0m \u001b[0;32melse\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 358\u001b[0m \u001b[0;32mif\u001b[0m \u001b[0mtimeout\u001b[0m \u001b[0;34m>\u001b[0m \u001b[0;36m0\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 359\u001b[0;31m \u001b[0mgotit\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mwaiter\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0macquire\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;32mTrue\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mtimeout\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m 360\u001b[0m \u001b[0;32melse\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 361\u001b[0m \u001b[0mgotit\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mwaiter\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0macquire\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;32mFalse\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n", - "\u001b[0;31mKeyboardInterrupt\u001b[0m: " - ] - } - ], + "outputs": [], "source": [ "import os\n", "os.environ.setdefault(\"WL_PRELOAD_IMAGE_OVERVIEW\", \"0\")\n", @@ -550,7 +155,6 @@ "import warnings\n", "warnings.filterwarnings(\"ignore\")\n", "\n", - "from weightslab.integrations.ultralytics import WLAwareTrainer\n", "from ultralytics import YOLO\n", "\n", "# Hyperparameters registered with WeightsLab -> live-editable from Weights Studio.\n", @@ -575,7 +179,7 @@ "# Read back the (now live) hyperparameters and hand training to WLAwareTrainer.\n", "model = YOLO(cfg[\"model\"])\n", "results = model.train(\n", - " trainer=WLAwareTrainer,\n", + " trainer=wl.WLAwareTrainer,\n", " data=str(cfg[\"data\"]),\n", " imgsz=cfg[\"image_size\"],\n", " epochs=int(cfg[\"epochs\"]),\n", @@ -606,7 +210,7 @@ }, { "cell_type": "code", - "execution_count": 3, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/", @@ -616,1230 +220,7 @@ "id": "uOJdXGA3Ks7o", "outputId": "26652fa3-6a15-43ce-9d75-034eaba21ccd" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "1241 samples tracked\n" - ] - }, - { - "data": { - "application/vnd.google.colaboratory.intrinsic+json": { - "repr_error": "Out of range float values are not JSON compliant: nan", - "type": "dataframe" - }, - "text/html": [ - "\n", - "
\n", - "
\n", - "\n", - "\n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - "
predictionprediction_rawtargetdiscardedorigintask_typelast_seengroup_idmember_rankimg_pathcls
sample_idannotation_id
00NoneNone[0.2335685, 0.254695, 0.4553995, 0.43075103, 1...Falsetrain_loader-1.000.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
10NoneNone[0.251174, 0.24061048, 0.44366202, 0.4307515, ...Falsetrain_loader-1.010.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
20NoneNoneNoneFalsetrain_loader-1.020.0/content/datasets/brain-tumor/images/train/000...[[1.0], [1.0]]
1NoneNone[0.523474, 0.3978875, 0.634976, 0.48122054, 1....NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.40610352, 0.415493, 0.5234745, 0.522301, 1....NoneNaNNoneNaNNoneNaNNoneNone
30NoneNone[0.403756, 0.3826295, 0.637324, 0.5152585, 1.0...Falsetrain_loader-1.030.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
40NoneNone[0.4436615, 0.3814555, 0.5938965, 0.4518785, 1...Falsetrain_loader-1.040.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
50NoneNone[0.6044605, 0.3990615, 0.67253554, 0.46596247,...Falsetrain_loader-1.050.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
60NoneNone[0.410798, 0.4436625, 0.556338, 0.51173747, 1....Falsetrain_loader-1.060.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
70NoneNone[0.650235, 0.4166665, 0.73826295, 0.48356748, ...Falsetrain_loader-1.070.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
80NoneNone[0.64788747, 0.388498, 0.7582165, 0.49413198, ...Falsetrain_loader-1.080.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
90NoneNone[0.6807515, 0.40140802, 0.75704247, 0.45774597...Falsetrain_loader-1.090.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
100NoneNone[0.431925, 0.1701875, 0.561033, 0.3180745, 1.0...Falsetrain_loader-1.0100.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
110NoneNone[0.42018852, 0.34507048, 0.5434275, 0.4941315,...Falsetrain_loader-1.0110.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
120NoneNone[0.3732395, 0.348592, 0.54812247, 0.46009403, ...Falsetrain_loader-1.0120.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
130NoneNone[0.37793398, 0.3814555, 0.48122, 0.45892054, 1...Falsetrain_loader-1.0130.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
140NoneNoneNoneFalsetrain_loader-1.0140.0/content/datasets/brain-tumor/images/train/000...[[1.0], [1.0]]
1NoneNone[0.43427247, 0.3556335, 0.51173747, 0.43075046...NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.4929575, 0.44718304, 0.5375585, 0.485915, 1...NoneNaNNoneNaNNoneNaNNoneNone
150NoneNoneNoneFalsetrain_loader-1.0150.0/content/datasets/brain-tumor/images/train/000...[[1.0], [1.0]]
1NoneNone[0.629108, 0.39788702, 0.67840403, 0.453051, 1...NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.4530515, 0.4107985, 0.58568054, 0.5105635, ...NoneNaNNoneNaNNoneNaNNoneNone
160NoneNoneNoneFalsetrain_loader-1.0160.0/content/datasets/brain-tumor/images/train/000...[[1.0], [1.0]]
1NoneNone[0.45070451, 0.3967135, 0.5762915, 0.5152585, ...NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.63732404, 0.403756, 0.684272, 0.46009403, 1...NoneNaNNoneNaNNoneNaNNoneNone
170NoneNone[0.43779355, 0.3826295, 0.61032856, 0.5223005,...Falsetrain_loader-1.0170.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
180NoneNone[0.4401405, 0.374413, 0.6068075, 0.529343, 1.0...Falsetrain_loader-1.0180.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
190NoneNone[0.536385, 0.4495305, 0.600939, 0.49999946, 0....Falsetrain_loader-1.0190.0/content/datasets/brain-tumor/images/train/000...[[0.0]]
200NoneNoneNoneFalsetrain_loader-1.0200.0/content/datasets/brain-tumor/images/train/000...[[0.0], [0.0], [0.0]]
1NoneNone[0.52347445, 0.4424885, 0.6150235, 0.5281695, ...NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.456573, 0.515258, 0.504695, 0.562206, 0.0, ...NoneNaNNoneNaNNoneNaNNoneNone
3NoneNone[0.504695, 0.44248852, 0.54342705, 0.48004752,...NoneNaNNoneNaNNoneNaNNoneNone
210NoneNone[0.54107946, 0.456573, 0.6103285, 0.511737, 0....Falsetrain_loader-1.0210.0/content/datasets/brain-tumor/images/train/000...[[0.0]]
220NoneNone[0.3298125, 0.2042255, 0.4225355, 0.2957745, 0...Falsetrain_loader-1.0220.0/content/datasets/brain-tumor/images/train/000...[[0.0]]
230NoneNone[0.517606, 0.18661949, 0.615024, 0.2875585, 1....Falsetrain_loader-1.0230.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
240NoneNone[0.51173747, 0.18779299, 0.6126765, 0.308685, ...Falsetrain_loader-1.0240.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
250NoneNoneNoneFalsetrain_loader-1.0250.0/content/datasets/brain-tumor/images/train/000...[[0.0], [0.0]]
1NoneNone[0.30985948, 0.3626765, 0.3861505, 0.43896753,...NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.41314548, 0.3427225, 0.48122054, 0.42018753...NoneNaNNoneNaNNoneNaNNoneNone
260NoneNone[0.25821602, 0.30046952, 0.46830997, 0.4577465...Falsetrain_loader-1.0260.0/content/datasets/brain-tumor/images/train/000...[[0.0]]
270NoneNone[0.2734745, 0.2828635, 0.44718352, 0.4577465, ...Falsetrain_loader-1.0270.0/content/datasets/brain-tumor/images/train/000...[[0.0]]
280NoneNone[0.5915495, 0.23591551, 0.6455405, 0.2887325, ...Falsetrain_loader-1.0280.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
290NoneNone[0.5903755, 0.231221, 0.6420185, 0.280517, 1.0...Falsetrain_loader-1.0290.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
300NoneNone[0.40493003, 0.2018775, 0.469484, 0.2699525, 1...Falsetrain_loader-1.0300.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
310NoneNone[0.39554, 0.19483551, 0.483568, 0.2887325, 1.0...Falsetrain_loader-1.0310.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
320NoneNone[0.4025825, 0.212441, 0.45539945, 0.269953, 1....Falsetrain_loader-1.0320.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
330NoneNone[0.2265255, 0.2136145, 0.3438965, 0.30751148, ...Falsetrain_loader-1.0330.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
340NoneNone[0.634976, 0.349765, 0.7723, 0.456573, 1.0, 1.0]Falsetrain_loader-1.0340.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
350NoneNone[0.6079815, 0.320423, 0.7969485, 0.489437, 1.0...Falsetrain_loader-1.0350.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
360NoneNone[0.6185445, 0.30164298, 0.7758215, 0.494131, 1...Falsetrain_loader-1.0360.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
\n", - "
\n", - "
\n", - "\n", - "
\n", - " \n", - "\n", - " \n", - "\n", - " \n", - "
\n", - "\n", - "\n", - "
\n", - "
\n" - ], - "text/plain": [ - " prediction prediction_raw \\\n", - "sample_id annotation_id \n", - "0 0 None None \n", - "1 0 None None \n", - "2 0 None None \n", - " 1 None None \n", - " 2 None None \n", - "3 0 None None \n", - "4 0 None None \n", - "5 0 None None \n", - "6 0 None None \n", - "7 0 None None \n", - "8 0 None None \n", - "9 0 None None \n", - "10 0 None None \n", - "11 0 None None \n", - "12 0 None None \n", - "13 0 None None \n", - "14 0 None None \n", - " 1 None None \n", - " 2 None None \n", - "15 0 None None \n", - " 1 None None \n", - " 2 None None \n", - "16 0 None None \n", - " 1 None None \n", - " 2 None None \n", - "17 0 None None \n", - "18 0 None None \n", - "19 0 None None \n", - "20 0 None None \n", - " 1 None None \n", - " 2 None None \n", - " 3 None None \n", - "21 0 None None \n", - "22 0 None None \n", - "23 0 None None \n", - "24 0 None None \n", - "25 0 None None \n", - " 1 None None \n", - " 2 None None \n", - "26 0 None None \n", - "27 0 None None \n", - "28 0 None None \n", - "29 0 None None \n", - "30 0 None None \n", - "31 0 None None \n", - "32 0 None None \n", - "33 0 None None \n", - "34 0 None None \n", - "35 0 None None \n", - "36 0 None None \n", - "\n", - " target \\\n", - "sample_id annotation_id \n", - "0 0 [0.2335685, 0.254695, 0.4553995, 0.43075103, 1... \n", - "1 0 [0.251174, 0.24061048, 0.44366202, 0.4307515, ... \n", - "2 0 None \n", - " 1 [0.523474, 0.3978875, 0.634976, 0.48122054, 1.... \n", - " 2 [0.40610352, 0.415493, 0.5234745, 0.522301, 1.... \n", - "3 0 [0.403756, 0.3826295, 0.637324, 0.5152585, 1.0... \n", - "4 0 [0.4436615, 0.3814555, 0.5938965, 0.4518785, 1... \n", - "5 0 [0.6044605, 0.3990615, 0.67253554, 0.46596247,... \n", - "6 0 [0.410798, 0.4436625, 0.556338, 0.51173747, 1.... \n", - "7 0 [0.650235, 0.4166665, 0.73826295, 0.48356748, ... \n", - "8 0 [0.64788747, 0.388498, 0.7582165, 0.49413198, ... \n", - "9 0 [0.6807515, 0.40140802, 0.75704247, 0.45774597... \n", - "10 0 [0.431925, 0.1701875, 0.561033, 0.3180745, 1.0... \n", - "11 0 [0.42018852, 0.34507048, 0.5434275, 0.4941315,... \n", - "12 0 [0.3732395, 0.348592, 0.54812247, 0.46009403, ... \n", - "13 0 [0.37793398, 0.3814555, 0.48122, 0.45892054, 1... \n", - "14 0 None \n", - " 1 [0.43427247, 0.3556335, 0.51173747, 0.43075046... \n", - " 2 [0.4929575, 0.44718304, 0.5375585, 0.485915, 1... \n", - "15 0 None \n", - " 1 [0.629108, 0.39788702, 0.67840403, 0.453051, 1... \n", - " 2 [0.4530515, 0.4107985, 0.58568054, 0.5105635, ... \n", - "16 0 None \n", - " 1 [0.45070451, 0.3967135, 0.5762915, 0.5152585, ... \n", - " 2 [0.63732404, 0.403756, 0.684272, 0.46009403, 1... \n", - "17 0 [0.43779355, 0.3826295, 0.61032856, 0.5223005,... \n", - "18 0 [0.4401405, 0.374413, 0.6068075, 0.529343, 1.0... \n", - "19 0 [0.536385, 0.4495305, 0.600939, 0.49999946, 0.... \n", - "20 0 None \n", - " 1 [0.52347445, 0.4424885, 0.6150235, 0.5281695, ... \n", - " 2 [0.456573, 0.515258, 0.504695, 0.562206, 0.0, ... \n", - " 3 [0.504695, 0.44248852, 0.54342705, 0.48004752,... \n", - "21 0 [0.54107946, 0.456573, 0.6103285, 0.511737, 0.... \n", - "22 0 [0.3298125, 0.2042255, 0.4225355, 0.2957745, 0... \n", - "23 0 [0.517606, 0.18661949, 0.615024, 0.2875585, 1.... \n", - "24 0 [0.51173747, 0.18779299, 0.6126765, 0.308685, ... \n", - "25 0 None \n", - " 1 [0.30985948, 0.3626765, 0.3861505, 0.43896753,... \n", - " 2 [0.41314548, 0.3427225, 0.48122054, 0.42018753... \n", - "26 0 [0.25821602, 0.30046952, 0.46830997, 0.4577465... \n", - "27 0 [0.2734745, 0.2828635, 0.44718352, 0.4577465, ... \n", - "28 0 [0.5915495, 0.23591551, 0.6455405, 0.2887325, ... \n", - "29 0 [0.5903755, 0.231221, 0.6420185, 0.280517, 1.0... \n", - "30 0 [0.40493003, 0.2018775, 0.469484, 0.2699525, 1... \n", - "31 0 [0.39554, 0.19483551, 0.483568, 0.2887325, 1.0... \n", - "32 0 [0.4025825, 0.212441, 0.45539945, 0.269953, 1.... \n", - "33 0 [0.2265255, 0.2136145, 0.3438965, 0.30751148, ... \n", - "34 0 [0.634976, 0.349765, 0.7723, 0.456573, 1.0, 1.0] \n", - "35 0 [0.6079815, 0.320423, 0.7969485, 0.489437, 1.0... \n", - "36 0 [0.6185445, 0.30164298, 0.7758215, 0.494131, 1... \n", - "\n", - " discarded origin task_type last_seen group_id \\\n", - "sample_id annotation_id \n", - "0 0 False train_loader -1.0 0 \n", - "1 0 False train_loader -1.0 1 \n", - "2 0 False train_loader -1.0 2 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - "3 0 False train_loader -1.0 3 \n", - "4 0 False train_loader -1.0 4 \n", - "5 0 False train_loader -1.0 5 \n", - "6 0 False train_loader -1.0 6 \n", - "7 0 False train_loader -1.0 7 \n", - "8 0 False train_loader -1.0 8 \n", - "9 0 False train_loader -1.0 9 \n", - "10 0 False train_loader -1.0 10 \n", - "11 0 False train_loader -1.0 11 \n", - "12 0 False train_loader -1.0 12 \n", - "13 0 False train_loader -1.0 13 \n", - "14 0 False train_loader -1.0 14 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - "15 0 False train_loader -1.0 15 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - "16 0 False train_loader -1.0 16 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - "17 0 False train_loader -1.0 17 \n", - "18 0 False train_loader -1.0 18 \n", - "19 0 False train_loader -1.0 19 \n", - "20 0 False train_loader -1.0 20 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - " 3 None NaN None NaN None \n", - "21 0 False train_loader -1.0 21 \n", - "22 0 False train_loader -1.0 22 \n", - "23 0 False train_loader -1.0 23 \n", - "24 0 False train_loader -1.0 24 \n", - "25 0 False train_loader -1.0 25 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - "26 0 False train_loader -1.0 26 \n", - "27 0 False train_loader -1.0 27 \n", - "28 0 False train_loader -1.0 28 \n", - "29 0 False train_loader -1.0 29 \n", - "30 0 False train_loader -1.0 30 \n", - "31 0 False train_loader -1.0 31 \n", - "32 0 False train_loader -1.0 32 \n", - "33 0 False train_loader -1.0 33 \n", - "34 0 False train_loader -1.0 34 \n", - "35 0 False train_loader -1.0 35 \n", - "36 0 False train_loader -1.0 36 \n", - "\n", - " member_rank \\\n", - "sample_id annotation_id \n", - "0 0 0.0 \n", - "1 0 0.0 \n", - "2 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - "3 0 0.0 \n", - "4 0 0.0 \n", - "5 0 0.0 \n", - "6 0 0.0 \n", - "7 0 0.0 \n", - "8 0 0.0 \n", - "9 0 0.0 \n", - "10 0 0.0 \n", - "11 0 0.0 \n", - "12 0 0.0 \n", - "13 0 0.0 \n", - "14 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - "15 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - "16 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - "17 0 0.0 \n", - "18 0 0.0 \n", - "19 0 0.0 \n", - "20 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - " 3 NaN \n", - "21 0 0.0 \n", - "22 0 0.0 \n", - "23 0 0.0 \n", - "24 0 0.0 \n", - "25 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - "26 0 0.0 \n", - "27 0 0.0 \n", - "28 0 0.0 \n", - "29 0 0.0 \n", - "30 0 0.0 \n", - "31 0 0.0 \n", - "32 0 0.0 \n", - "33 0 0.0 \n", - "34 0 0.0 \n", - "35 0 0.0 \n", - "36 0 0.0 \n", - "\n", - " img_path \\\n", - "sample_id annotation_id \n", - "0 0 /content/datasets/brain-tumor/images/train/000... \n", - "1 0 /content/datasets/brain-tumor/images/train/000... \n", - "2 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - "3 0 /content/datasets/brain-tumor/images/train/000... \n", - "4 0 /content/datasets/brain-tumor/images/train/000... \n", - "5 0 /content/datasets/brain-tumor/images/train/000... \n", - "6 0 /content/datasets/brain-tumor/images/train/000... \n", - "7 0 /content/datasets/brain-tumor/images/train/000... \n", - "8 0 /content/datasets/brain-tumor/images/train/000... \n", - "9 0 /content/datasets/brain-tumor/images/train/000... \n", - "10 0 /content/datasets/brain-tumor/images/train/000... \n", - "11 0 /content/datasets/brain-tumor/images/train/000... \n", - "12 0 /content/datasets/brain-tumor/images/train/000... \n", - "13 0 /content/datasets/brain-tumor/images/train/000... \n", - "14 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - "15 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - "16 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - "17 0 /content/datasets/brain-tumor/images/train/000... \n", - "18 0 /content/datasets/brain-tumor/images/train/000... \n", - "19 0 /content/datasets/brain-tumor/images/train/000... \n", - "20 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - " 3 None \n", - "21 0 /content/datasets/brain-tumor/images/train/000... \n", - "22 0 /content/datasets/brain-tumor/images/train/000... \n", - "23 0 /content/datasets/brain-tumor/images/train/000... \n", - "24 0 /content/datasets/brain-tumor/images/train/000... \n", - "25 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - "26 0 /content/datasets/brain-tumor/images/train/000... \n", - "27 0 /content/datasets/brain-tumor/images/train/000... \n", - "28 0 /content/datasets/brain-tumor/images/train/000... \n", - "29 0 /content/datasets/brain-tumor/images/train/000... \n", - "30 0 /content/datasets/brain-tumor/images/train/000... \n", - "31 0 /content/datasets/brain-tumor/images/train/000... \n", - "32 0 /content/datasets/brain-tumor/images/train/000... \n", - "33 0 /content/datasets/brain-tumor/images/train/000... \n", - "34 0 /content/datasets/brain-tumor/images/train/000... \n", - "35 0 /content/datasets/brain-tumor/images/train/000... \n", - "36 0 /content/datasets/brain-tumor/images/train/000... \n", - "\n", - " cls \n", - "sample_id annotation_id \n", - "0 0 [[1.0]] \n", - "1 0 [[1.0]] \n", - "2 0 [[1.0], [1.0]] \n", - " 1 None \n", - " 2 None \n", - "3 0 [[1.0]] \n", - "4 0 [[1.0]] \n", - "5 0 [[1.0]] \n", - "6 0 [[1.0]] \n", - "7 0 [[1.0]] \n", - "8 0 [[1.0]] \n", - "9 0 [[1.0]] \n", - "10 0 [[1.0]] \n", - "11 0 [[1.0]] \n", - "12 0 [[1.0]] \n", - "13 0 [[1.0]] \n", - "14 0 [[1.0], [1.0]] \n", - " 1 None \n", - " 2 None \n", - "15 0 [[1.0], [1.0]] \n", - " 1 None \n", - " 2 None \n", - "16 0 [[1.0], [1.0]] \n", - " 1 None \n", - " 2 None \n", - "17 0 [[1.0]] \n", - "18 0 [[1.0]] \n", - "19 0 [[0.0]] \n", - "20 0 [[0.0], [0.0], [0.0]] \n", - " 1 None \n", - " 2 None \n", - " 3 None \n", - "21 0 [[0.0]] \n", - "22 0 [[0.0]] \n", - "23 0 [[1.0]] \n", - "24 0 [[1.0]] \n", - "25 0 [[0.0], [0.0]] \n", - " 1 None \n", - " 2 None \n", - "26 0 [[0.0]] \n", - "27 0 [[0.0]] \n", - "28 0 [[1.0]] \n", - "29 0 [[1.0]] \n", - "30 0 [[1.0]] \n", - "31 0 [[1.0]] \n", - "32 0 [[1.0]] \n", - "33 0 [[1.0]] \n", - "34 0 [[1.0]] \n", - "35 0 [[1.0]] \n", - "36 0 [[1.0]] " - ] - }, - "metadata": {}, - "output_type": "display_data" - }, - { - "data": { - "text/html": [ - "🔗 Open the full grid in Weights Studio" - ], - "text/plain": [ - "" - ] - }, - "metadata": {}, - "output_type": "display_data" - } - ], + "outputs": [], "source": [ "import pandas as pd\n", "from IPython.display import display, HTML\n", diff --git a/weightslab/examples/Notebooks/Ultralytics/wl-how-to-train-ultralytics-yolo-on-carparts-segmentation-dataset.ipynb b/weightslab/examples/Notebooks/Ultralytics/wl-how-to-train-ultralytics-yolo-on-carparts-segmentation-dataset.ipynb index 93fca215..99c57f0a 100644 --- a/weightslab/examples/Notebooks/Ultralytics/wl-how-to-train-ultralytics-yolo-on-carparts-segmentation-dataset.ipynb +++ b/weightslab/examples/Notebooks/Ultralytics/wl-how-to-train-ultralytics-yolo-on-carparts-segmentation-dataset.ipynb @@ -44,7 +44,7 @@ }, { "cell_type": "code", - "execution_count": 6, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/" @@ -52,24 +52,14 @@ "id": "71be394c", "outputId": "c9ef4767-84f5-4bea-e395-9e11766002f2" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Requirement already satisfied: setuptools in /usr/local/lib/python3.12/dist-packages (75.2.0)\n", - "Requirement already satisfied: wheel in /usr/local/lib/python3.12/dist-packages (0.47.0)\n", - "Requirement already satisfied: packaging>=24.0 in /usr/local/lib/python3.12/dist-packages (from wheel) (26.2)\n" - ] - } - ], + "outputs": [], "source": [ "%pip install setuptools wheel" ] }, { "cell_type": "code", - "execution_count": 7, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/", @@ -78,149 +68,7 @@ "id": "jV1oi_PQN0xK", "outputId": "eda8eca1-f4b3-4ffd-eaa6-0f67ce857153" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Collecting git+https://github.com/GrayboxTech/weightslab.git@landingcollab\n", - " Cloning https://github.com/GrayboxTech/weightslab.git (to revision landingcollab) to /tmp/pip-req-build-z4ba7cx9\n", - " Running command git clone --filter=blob:none --quiet https://github.com/GrayboxTech/weightslab.git /tmp/pip-req-build-z4ba7cx9\n", - " Running command git checkout -b landingcollab --track origin/landingcollab\n", - " Switched to a new branch 'landingcollab'\n", - " Branch 'landingcollab' set up to track remote branch 'landingcollab' from 'origin'.\n", - " Resolved https://github.com/GrayboxTech/weightslab.git to commit 3f1cf143bc93e67df43adc71ab831fa71477e0d2\n", - " Installing build dependencies ... \u001b[?25l\u001b[?25hdone\n", - " Getting requirements to build wheel ... \u001b[?25l\u001b[?25hdone\n", - " Preparing metadata (pyproject.toml) ... \u001b[?25l\u001b[?25hdone\n", - "Requirement already satisfied: numpy<3,>=1.24 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (2.0.2)\n", - "Requirement already satisfied: pandas<3,>=2.2.2 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (2.2.2)\n", - "Requirement already satisfied: duckdb<2,>=1.1 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.3.2)\n", - "Requirement already satisfied: PyYAML<7,>=6.0.3 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (6.0.3)\n", - "Requirement already satisfied: dill<0.5,>=0.3.8 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (0.3.8)\n", - "Requirement already satisfied: zstandard<1,>=0.22 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (0.25.0)\n", - "Requirement already satisfied: h5py<4,>=3.10 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (3.16.0)\n", - "Requirement already satisfied: xxhash<4.1,>=3.4 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (3.7.1)\n", - "Requirement already satisfied: tables<4,>=3.9 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (3.10.2)\n", - "Requirement already satisfied: torch<=2.9,>=2.1 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (2.9.0)\n", - "Requirement already satisfied: torchvision<1,>=0.16 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (0.24.0)\n", - "Requirement already satisfied: torchmetrics>=1.9 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.9.0)\n", - "Requirement already satisfied: grpcio<2,>=1.80 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.81.1)\n", - "Requirement already satisfied: protobuf<8,>=5.28.1 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (5.29.6)\n", - "Requirement already satisfied: pydantic<3,>=2.7 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (2.13.4)\n", - "Requirement already satisfied: Pillow<12,>=10 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (11.3.0)\n", - "Requirement already satisfied: graphviz<1,>=0.20 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (0.21)\n", - "Requirement already satisfied: onnx<=1.20,>=1.15 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.20.0)\n", - "Requirement already satisfied: tqdm<5,>=4.66 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (4.67.3)\n", - "Requirement already satisfied: python-dotenv<2,>=1 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.2.2)\n", - "Requirement already satisfied: langchain-core<2,>=0.3 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.4.9)\n", - "Requirement already satisfied: langchain-ollama<2,>=0.2 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.1.0)\n", - "Requirement already satisfied: langchain-openai<2,>=0.2 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.3.5)\n", - "Requirement already satisfied: typing-extensions~=4.12 in /usr/local/lib/python3.12/dist-packages (from grpcio<2,>=1.80->weightslab==1.3.3.dev8) (4.15.0)\n", - "Requirement already satisfied: jsonpatch<2.0.0,>=1.33.0 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (1.33)\n", - "Requirement already satisfied: langchain-protocol>=0.0.17 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (0.0.18)\n", - "Requirement already satisfied: langsmith<1.0.0,>=0.3.45 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (0.9.1)\n", - "Requirement already satisfied: packaging>=23.2.0 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (26.2)\n", - "Requirement already satisfied: tenacity!=8.4.0,<10.0.0,>=8.1.0 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (9.1.4)\n", - "Requirement already satisfied: uuid-utils<1.0,>=0.12.0 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (0.16.2)\n", - "Requirement already satisfied: ollama<1.0.0,>=0.6.1 in /usr/local/lib/python3.12/dist-packages (from langchain-ollama<2,>=0.2->weightslab==1.3.3.dev8) (0.6.2)\n", - "Requirement already satisfied: openai<3.0.0,>=2.45.0 in /usr/local/lib/python3.12/dist-packages (from langchain-openai<2,>=0.2->weightslab==1.3.3.dev8) (2.45.0)\n", - "Requirement already satisfied: tiktoken<1.0.0,>=0.7.0 in /usr/local/lib/python3.12/dist-packages (from langchain-openai<2,>=0.2->weightslab==1.3.3.dev8) (0.13.0)\n", - "Requirement already satisfied: ml_dtypes>=0.5.0 in /usr/local/lib/python3.12/dist-packages (from onnx<=1.20,>=1.15->weightslab==1.3.3.dev8) (0.5.4)\n", - "Requirement already satisfied: python-dateutil>=2.8.2 in /usr/local/lib/python3.12/dist-packages (from pandas<3,>=2.2.2->weightslab==1.3.3.dev8) (2.9.0.post0)\n", - "Requirement already satisfied: pytz>=2020.1 in /usr/local/lib/python3.12/dist-packages (from pandas<3,>=2.2.2->weightslab==1.3.3.dev8) (2025.2)\n", - "Requirement already satisfied: tzdata>=2022.7 in /usr/local/lib/python3.12/dist-packages (from pandas<3,>=2.2.2->weightslab==1.3.3.dev8) (2026.2)\n", - "Requirement already satisfied: annotated-types>=0.6.0 in /usr/local/lib/python3.12/dist-packages (from pydantic<3,>=2.7->weightslab==1.3.3.dev8) (0.7.0)\n", - "Requirement already satisfied: pydantic-core==2.46.4 in /usr/local/lib/python3.12/dist-packages (from pydantic<3,>=2.7->weightslab==1.3.3.dev8) (2.46.4)\n", - "Requirement already satisfied: typing-inspection>=0.4.2 in /usr/local/lib/python3.12/dist-packages (from pydantic<3,>=2.7->weightslab==1.3.3.dev8) (0.4.2)\n", - "Requirement already satisfied: numexpr>=2.6.2 in /usr/local/lib/python3.12/dist-packages (from tables<4,>=3.9->weightslab==1.3.3.dev8) (2.14.1)\n", - "Requirement already satisfied: py-cpuinfo in /usr/local/lib/python3.12/dist-packages (from tables<4,>=3.9->weightslab==1.3.3.dev8) (9.0.0)\n", - "Requirement already satisfied: blosc2>=2.3.0 in /usr/local/lib/python3.12/dist-packages (from tables<4,>=3.9->weightslab==1.3.3.dev8) (4.5.1)\n", - "Requirement already satisfied: filelock in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.29.4)\n", - "Requirement already satisfied: setuptools in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (75.2.0)\n", - "Requirement already satisfied: sympy>=1.13.3 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (1.14.0)\n", - "Requirement already satisfied: networkx>=2.5.1 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.6.1)\n", - "Requirement already satisfied: jinja2 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.1.6)\n", - "Requirement already satisfied: fsspec>=0.8.5 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (2025.3.0)\n", - "Requirement already satisfied: nvidia-cuda-nvrtc-cu12==12.8.93 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.93)\n", - "Requirement already satisfied: nvidia-cuda-runtime-cu12==12.8.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.90)\n", - "Requirement already satisfied: nvidia-cuda-cupti-cu12==12.8.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.90)\n", - "Requirement already satisfied: nvidia-cudnn-cu12==9.10.2.21 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (9.10.2.21)\n", - "Requirement already satisfied: nvidia-cublas-cu12==12.8.4.1 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.4.1)\n", - "Requirement already satisfied: nvidia-cufft-cu12==11.3.3.83 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (11.3.3.83)\n", - "Requirement already satisfied: nvidia-curand-cu12==10.3.9.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (10.3.9.90)\n", - "Requirement already satisfied: nvidia-cusolver-cu12==11.7.3.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (11.7.3.90)\n", - "Requirement already satisfied: nvidia-cusparse-cu12==12.5.8.93 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.5.8.93)\n", - "Requirement already satisfied: nvidia-cusparselt-cu12==0.7.1 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (0.7.1)\n", - "Requirement already satisfied: nvidia-nccl-cu12==2.27.5 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (2.27.5)\n", - "Requirement already satisfied: nvidia-nvshmem-cu12==3.3.20 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.3.20)\n", - "Requirement already satisfied: nvidia-nvtx-cu12==12.8.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.90)\n", - "Requirement already satisfied: nvidia-nvjitlink-cu12==12.8.93 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.93)\n", - "Requirement already satisfied: nvidia-cufile-cu12==1.13.1.3 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (1.13.1.3)\n", - "Requirement already satisfied: triton==3.5.0 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.5.0)\n", - "Requirement already satisfied: lightning-utilities>=0.15.3 in /usr/local/lib/python3.12/dist-packages (from torchmetrics>=1.9->weightslab==1.3.3.dev8) (0.15.3)\n", - "Requirement already satisfied: ndindex in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (1.10.1)\n", - "Requirement already satisfied: msgpack in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (1.2.1)\n", - "Requirement already satisfied: requests in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (2.32.4)\n", - "Requirement already satisfied: rich in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (13.9.4)\n", - "Requirement already satisfied: textual in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (6.2.1)\n", - "Requirement already satisfied: textual-plotext in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (1.0.1)\n", - "Requirement already satisfied: threadpoolctl in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (3.6.0)\n", - "Requirement already satisfied: jsonpointer>=1.9 in /usr/local/lib/python3.12/dist-packages (from jsonpatch<2.0.0,>=1.33.0->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (3.1.1)\n", - "Requirement already satisfied: anyio>=3.5.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (4.14.0)\n", - "Requirement already satisfied: distro>=1.7.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (1.9.0)\n", - "Requirement already satisfied: httpx<1,>=0.23.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (0.28.1)\n", - "Requirement already satisfied: orjson>=3.9.14 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (3.11.9)\n", - "Requirement already satisfied: requests-toolbelt>=1.0.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (1.0.0)\n", - "Requirement already satisfied: sniffio>=1.1 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (1.3.1)\n", - "Requirement already satisfied: websockets>=15.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (15.0.1)\n", - "Requirement already satisfied: jiter<1,>=0.10.0 in /usr/local/lib/python3.12/dist-packages (from openai<3.0.0,>=2.45.0->langchain-openai<2,>=0.2->weightslab==1.3.3.dev8) (0.15.0)\n", - "Requirement already satisfied: six>=1.5 in /usr/local/lib/python3.12/dist-packages (from python-dateutil>=2.8.2->pandas<3,>=2.2.2->weightslab==1.3.3.dev8) (1.17.0)\n", - "Requirement already satisfied: mpmath<1.4,>=1.1.0 in /usr/local/lib/python3.12/dist-packages (from sympy>=1.13.3->torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (1.3.0)\n", - "Requirement already satisfied: regex in /usr/local/lib/python3.12/dist-packages (from tiktoken<1.0.0,>=0.7.0->langchain-openai<2,>=0.2->weightslab==1.3.3.dev8) (2025.11.3)\n", - "Requirement already satisfied: MarkupSafe>=2.0 in /usr/local/lib/python3.12/dist-packages (from jinja2->torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.0.3)\n", - "Requirement already satisfied: idna>=2.8 in /usr/local/lib/python3.12/dist-packages (from anyio>=3.5.0->langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (3.18)\n", - "Requirement already satisfied: certifi in /usr/local/lib/python3.12/dist-packages (from httpx<1,>=0.23.0->langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (2026.6.17)\n", - "Requirement already satisfied: httpcore==1.* in /usr/local/lib/python3.12/dist-packages (from httpx<1,>=0.23.0->langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (1.0.9)\n", - "Requirement already satisfied: h11>=0.16 in /usr/local/lib/python3.12/dist-packages (from httpcore==1.*->httpx<1,>=0.23.0->langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (0.16.0)\n", - "Requirement already satisfied: charset_normalizer<4,>=2 in /usr/local/lib/python3.12/dist-packages (from requests->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (3.4.7)\n", - "Requirement already satisfied: urllib3<3,>=1.21.1 in /usr/local/lib/python3.12/dist-packages (from requests->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (2.5.0)\n", - "Requirement already satisfied: markdown-it-py>=2.2.0 in /usr/local/lib/python3.12/dist-packages (from rich->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (4.2.0)\n", - "Requirement already satisfied: pygments<3.0.0,>=2.13.0 in /usr/local/lib/python3.12/dist-packages (from rich->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (2.20.0)\n", - "Requirement already satisfied: platformdirs<5,>=3.6.0 in /usr/local/lib/python3.12/dist-packages (from textual->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (4.10.0)\n", - "Requirement already satisfied: plotext<6.0.0,>=5.2.8 in /usr/local/lib/python3.12/dist-packages (from textual-plotext->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (5.3.2)\n", - "Requirement already satisfied: mdurl~=0.1 in /usr/local/lib/python3.12/dist-packages (from markdown-it-py>=2.2.0->rich->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (0.1.2)\n", - "Requirement already satisfied: linkify-it-py<3,>=1 in /usr/local/lib/python3.12/dist-packages (from markdown-it-py[linkify,plugins]>=2.1.0->textual->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (2.1.0)\n", - "Requirement already satisfied: mdit-py-plugins>=0.5.0 in /usr/local/lib/python3.12/dist-packages (from markdown-it-py[linkify,plugins]>=2.1.0->textual->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (0.6.1)\n", - "Requirement already satisfied: uc-micro-py in /usr/local/lib/python3.12/dist-packages (from linkify-it-py<3,>=1->markdown-it-py[linkify,plugins]>=2.1.0->textual->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (2.0.0)\n", - "Building wheels for collected packages: weightslab\n", - " Building wheel for weightslab (pyproject.toml) ... \u001b[?25l\u001b[?25hdone\n", - " Created wheel for weightslab: filename=weightslab-1.3.3.dev8-py3-none-any.whl size=2827715 sha256=5c4e7c14bbf95a48b9a8c86b5a8c744af67e1eb395197cbe79d3e443c21cf86e\n", - " Stored in directory: /tmp/pip-ephem-wheel-cache-__zwgf7x/wheels/4d/c7/cd/817c2745790555e7b6c532f5b6055d47567f9416fdfc65879b\n", - "Successfully built weightslab\n", - "Installing collected packages: weightslab\n", - " Attempting uninstall: weightslab\n", - " Found existing installation: weightslab 1.3.3.dev7\n", - " Uninstalling weightslab-1.3.3.dev7:\n", - " Successfully uninstalled weightslab-1.3.3.dev7\n", - "Successfully installed weightslab-1.3.3.dev8\n" - ] - }, - { - "data": { - "application/vnd.colab-display-data+json": { - "id": "b7fce274042740d19087972fe9c01c9b", - "pip_warning": { - "packages": [ - "weightslab" - ] - } - } - }, - "metadata": {}, - "output_type": "display_data" - } - ], + "outputs": [], "source": [ "%pip install weightslab\n", "!pip install --upgrade \"protobuf>=6.31.1\" # Google Collab compat update" @@ -228,7 +76,7 @@ }, { "cell_type": "code", - "execution_count": 1, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/" @@ -236,39 +84,7 @@ "id": "pJhvREbaKs7f", "outputId": "651baa36-9ebb-49da-830b-e6c8a6e2b64f" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Ultralytics 8.4.96 🚀 Python-3.12.13 torch-2.9.0+cu128 CPU (Intel Xeon CPU @ 2.20GHz)\n", - "Setup complete ✅ (2 CPUs, 12.7 GB RAM, 30.5/107.7 GB disk)\n", - "16/07/2026-09:29:29.683 INFO:root:setup_logging: WeightsLab logging initialized - Log file: /tmp/tmp6meuzqdz/weightslab_logs/weightslab_20260716_092929.log\n", - "16/07/2026-09:29:29.685 INFO:weightslab:: WeightsLab package initialized - Log level: INFO, Log to file: True\n", - "16/07/2026-09:29:29.686 INFO:weightslab:: \n", - "╭─────────────────────────────────────────────────────────────────────────────────────────────────────╮\n", - "│ │\n", - "│ \u001b[32m$$\\ $$\\ \u001b[0m $$\\ $$\\ $$\\ \u001b[31m$$\\ \u001b[0m $$\\ │\n", - "│ \u001b[32m$$ | $\\ $$ |\u001b[0m \\__| $$ | $$ | \u001b[31m$$ | \u001b[0m $$ | │\n", - "│ \u001b[32m$$ |$$$\\ $$ |\u001b[0m $$$$$$\\ $$\\ $$$$$$\\ $$$$$$$\\ $$$$$$\\ $$$$$$$\\ \u001b[31m$$ | \u001b[0m $$$$$$\\ $$$$$$$\\ │\n", - "│ \u001b[32m$$ $$ $$\\$$ |\u001b[0m$$ __$$\\ $$ |$$ __$$\\ $$ __$$\\ \\_$$ _| $$ _____|\u001b[31m$$ | \u001b[0m \\____$$\\ $$ __$$\\ │\n", - "│ \u001b[32m$$$$ _$$$$ |\u001b[0m$$$$$$$$ |$$ |$$ / $$ |$$ | $$ | $$ | \\$$$$$$\\ \u001b[31m$$ | \u001b[0m $$$$$$$ |$$ | $$ | │\n", - "│ \u001b[32m$$$ / \\$$$ |\u001b[0m$$ ____|$$ |$$ | $$ |$$ | $$ | $$ |$$\\ \\____$$\\ \u001b[31m$$ | \u001b[0m$$ __$$ |$$ | $$ | │\n", - "│ \u001b[32m$$ / \\$$ |\u001b[0m\\$$$$$$$\\ $$ |\\$$$$$$$ |$$ | $$ | \\$$$$ |$$$$$$$ |\u001b[31m$$$$$$$$\\ \u001b[0m\\$$$$$$$ |$$$$$$$ | │\n", - "│ \u001b[32m\\__/ \\__|\u001b[0m \\_______|\\__| \\____$$ |\\__| \\__| \\____/ \\_______/ \u001b[31m\\________|\u001b[0m \\_______|\\_______/ │\n", - "│ \u001b[32m \u001b[0m $$\\ $$ |\u001b[31m\u001b[0m │\n", - "│ \u001b[32m \u001b[0m \\$$$$$$ |\u001b[31m\u001b[0m │\n", - "│ \u001b[32m \u001b[0m \\______/\u001b[31m\u001b[0m │\n", - "│ │\n", - "│ Inspect - Edit - Evolve Neural Networks │\n", - "│ │\n", - "╰─── By GrayBx, v1.3.3.dev8 ──────────────────────────────────────────────────────────────────────────╯\n", - "\n", - "\n", - "16/07/2026-09:29:29.690 INFO:weightslab.utils.telemetry:ping_import: WeightsLab uses anonymous usage data (package version used, and OS name). Set WL_NO_TELEMETRY=1 to disable.\n" - ] - } - ], + "outputs": [], "source": [ "# Install WeightsLab. Colab already ships torch, torchvision, numpy, scikit-learn\n", "# and Pillow, so nothing extra is needed here.\n", @@ -372,7 +188,6 @@ "import warnings\n", "warnings.filterwarnings(\"ignore\")\n", "\n", - "from weightslab.integrations.ultralytics import WLAwareSegmentationTrainer\n", "from ultralytics import YOLO\n", "\n", "# Hyperparameters registered with WeightsLab -> live-editable from Weights Studio.\n", @@ -397,7 +212,7 @@ "# Read back the (now live) hyperparameters and hand training to WLAwareSegmentationTrainer.\n", "model = YOLO(cfg[\"model\"])\n", "results = model.train(\n", - " trainer=WLAwareSegmentationTrainer,\n", + " trainer=wl.WLAwareSegmentationTrainer,\n", " data=str(cfg[\"data\"]),\n", " imgsz=cfg[\"image_size\"],\n", " epochs=int(cfg[\"epochs\"]),\n", @@ -427,7 +242,7 @@ }, { "cell_type": "code", - "execution_count": 3, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/", @@ -437,1230 +252,7 @@ "id": "uOJdXGA3Ks7o", "outputId": "26652fa3-6a15-43ce-9d75-034eaba21ccd" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "1241 samples tracked\n" - ] - }, - { - "data": { - "application/vnd.google.colaboratory.intrinsic+json": { - "repr_error": "Out of range float values are not JSON compliant: nan", - "type": "dataframe" - }, - "text/html": [ - "\n", - "
\n", - "
\n", - "\n", - "\n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - "
predictionprediction_rawtargetdiscardedorigintask_typelast_seengroup_idmember_rankimg_pathcls
sample_idannotation_id
00NoneNone[0.2335685, 0.254695, 0.4553995, 0.43075103, 1...Falsetrain_loader-1.000.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
10NoneNone[0.251174, 0.24061048, 0.44366202, 0.4307515, ...Falsetrain_loader-1.010.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
20NoneNoneNoneFalsetrain_loader-1.020.0/content/datasets/brain-tumor/images/train/000...[[1.0], [1.0]]
1NoneNone[0.523474, 0.3978875, 0.634976, 0.48122054, 1....NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.40610352, 0.415493, 0.5234745, 0.522301, 1....NoneNaNNoneNaNNoneNaNNoneNone
30NoneNone[0.403756, 0.3826295, 0.637324, 0.5152585, 1.0...Falsetrain_loader-1.030.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
40NoneNone[0.4436615, 0.3814555, 0.5938965, 0.4518785, 1...Falsetrain_loader-1.040.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
50NoneNone[0.6044605, 0.3990615, 0.67253554, 0.46596247,...Falsetrain_loader-1.050.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
60NoneNone[0.410798, 0.4436625, 0.556338, 0.51173747, 1....Falsetrain_loader-1.060.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
70NoneNone[0.650235, 0.4166665, 0.73826295, 0.48356748, ...Falsetrain_loader-1.070.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
80NoneNone[0.64788747, 0.388498, 0.7582165, 0.49413198, ...Falsetrain_loader-1.080.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
90NoneNone[0.6807515, 0.40140802, 0.75704247, 0.45774597...Falsetrain_loader-1.090.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
100NoneNone[0.431925, 0.1701875, 0.561033, 0.3180745, 1.0...Falsetrain_loader-1.0100.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
110NoneNone[0.42018852, 0.34507048, 0.5434275, 0.4941315,...Falsetrain_loader-1.0110.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
120NoneNone[0.3732395, 0.348592, 0.54812247, 0.46009403, ...Falsetrain_loader-1.0120.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
130NoneNone[0.37793398, 0.3814555, 0.48122, 0.45892054, 1...Falsetrain_loader-1.0130.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
140NoneNoneNoneFalsetrain_loader-1.0140.0/content/datasets/brain-tumor/images/train/000...[[1.0], [1.0]]
1NoneNone[0.43427247, 0.3556335, 0.51173747, 0.43075046...NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.4929575, 0.44718304, 0.5375585, 0.485915, 1...NoneNaNNoneNaNNoneNaNNoneNone
150NoneNoneNoneFalsetrain_loader-1.0150.0/content/datasets/brain-tumor/images/train/000...[[1.0], [1.0]]
1NoneNone[0.629108, 0.39788702, 0.67840403, 0.453051, 1...NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.4530515, 0.4107985, 0.58568054, 0.5105635, ...NoneNaNNoneNaNNoneNaNNoneNone
160NoneNoneNoneFalsetrain_loader-1.0160.0/content/datasets/brain-tumor/images/train/000...[[1.0], [1.0]]
1NoneNone[0.45070451, 0.3967135, 0.5762915, 0.5152585, ...NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.63732404, 0.403756, 0.684272, 0.46009403, 1...NoneNaNNoneNaNNoneNaNNoneNone
170NoneNone[0.43779355, 0.3826295, 0.61032856, 0.5223005,...Falsetrain_loader-1.0170.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
180NoneNone[0.4401405, 0.374413, 0.6068075, 0.529343, 1.0...Falsetrain_loader-1.0180.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
190NoneNone[0.536385, 0.4495305, 0.600939, 0.49999946, 0....Falsetrain_loader-1.0190.0/content/datasets/brain-tumor/images/train/000...[[0.0]]
200NoneNoneNoneFalsetrain_loader-1.0200.0/content/datasets/brain-tumor/images/train/000...[[0.0], [0.0], [0.0]]
1NoneNone[0.52347445, 0.4424885, 0.6150235, 0.5281695, ...NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.456573, 0.515258, 0.504695, 0.562206, 0.0, ...NoneNaNNoneNaNNoneNaNNoneNone
3NoneNone[0.504695, 0.44248852, 0.54342705, 0.48004752,...NoneNaNNoneNaNNoneNaNNoneNone
210NoneNone[0.54107946, 0.456573, 0.6103285, 0.511737, 0....Falsetrain_loader-1.0210.0/content/datasets/brain-tumor/images/train/000...[[0.0]]
220NoneNone[0.3298125, 0.2042255, 0.4225355, 0.2957745, 0...Falsetrain_loader-1.0220.0/content/datasets/brain-tumor/images/train/000...[[0.0]]
230NoneNone[0.517606, 0.18661949, 0.615024, 0.2875585, 1....Falsetrain_loader-1.0230.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
240NoneNone[0.51173747, 0.18779299, 0.6126765, 0.308685, ...Falsetrain_loader-1.0240.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
250NoneNoneNoneFalsetrain_loader-1.0250.0/content/datasets/brain-tumor/images/train/000...[[0.0], [0.0]]
1NoneNone[0.30985948, 0.3626765, 0.3861505, 0.43896753,...NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.41314548, 0.3427225, 0.48122054, 0.42018753...NoneNaNNoneNaNNoneNaNNoneNone
260NoneNone[0.25821602, 0.30046952, 0.46830997, 0.4577465...Falsetrain_loader-1.0260.0/content/datasets/brain-tumor/images/train/000...[[0.0]]
270NoneNone[0.2734745, 0.2828635, 0.44718352, 0.4577465, ...Falsetrain_loader-1.0270.0/content/datasets/brain-tumor/images/train/000...[[0.0]]
280NoneNone[0.5915495, 0.23591551, 0.6455405, 0.2887325, ...Falsetrain_loader-1.0280.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
290NoneNone[0.5903755, 0.231221, 0.6420185, 0.280517, 1.0...Falsetrain_loader-1.0290.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
300NoneNone[0.40493003, 0.2018775, 0.469484, 0.2699525, 1...Falsetrain_loader-1.0300.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
310NoneNone[0.39554, 0.19483551, 0.483568, 0.2887325, 1.0...Falsetrain_loader-1.0310.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
320NoneNone[0.4025825, 0.212441, 0.45539945, 0.269953, 1....Falsetrain_loader-1.0320.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
330NoneNone[0.2265255, 0.2136145, 0.3438965, 0.30751148, ...Falsetrain_loader-1.0330.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
340NoneNone[0.634976, 0.349765, 0.7723, 0.456573, 1.0, 1.0]Falsetrain_loader-1.0340.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
350NoneNone[0.6079815, 0.320423, 0.7969485, 0.489437, 1.0...Falsetrain_loader-1.0350.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
360NoneNone[0.6185445, 0.30164298, 0.7758215, 0.494131, 1...Falsetrain_loader-1.0360.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
\n", - "
\n", - "
\n", - "\n", - "
\n", - " \n", - "\n", - " \n", - "\n", - " \n", - "
\n", - "\n", - "\n", - "
\n", - "
\n" - ], - "text/plain": [ - " prediction prediction_raw \\\n", - "sample_id annotation_id \n", - "0 0 None None \n", - "1 0 None None \n", - "2 0 None None \n", - " 1 None None \n", - " 2 None None \n", - "3 0 None None \n", - "4 0 None None \n", - "5 0 None None \n", - "6 0 None None \n", - "7 0 None None \n", - "8 0 None None \n", - "9 0 None None \n", - "10 0 None None \n", - "11 0 None None \n", - "12 0 None None \n", - "13 0 None None \n", - "14 0 None None \n", - " 1 None None \n", - " 2 None None \n", - "15 0 None None \n", - " 1 None None \n", - " 2 None None \n", - "16 0 None None \n", - " 1 None None \n", - " 2 None None \n", - "17 0 None None \n", - "18 0 None None \n", - "19 0 None None \n", - "20 0 None None \n", - " 1 None None \n", - " 2 None None \n", - " 3 None None \n", - "21 0 None None \n", - "22 0 None None \n", - "23 0 None None \n", - "24 0 None None \n", - "25 0 None None \n", - " 1 None None \n", - " 2 None None \n", - "26 0 None None \n", - "27 0 None None \n", - "28 0 None None \n", - "29 0 None None \n", - "30 0 None None \n", - "31 0 None None \n", - "32 0 None None \n", - "33 0 None None \n", - "34 0 None None \n", - "35 0 None None \n", - "36 0 None None \n", - "\n", - " target \\\n", - "sample_id annotation_id \n", - "0 0 [0.2335685, 0.254695, 0.4553995, 0.43075103, 1... \n", - "1 0 [0.251174, 0.24061048, 0.44366202, 0.4307515, ... \n", - "2 0 None \n", - " 1 [0.523474, 0.3978875, 0.634976, 0.48122054, 1.... \n", - " 2 [0.40610352, 0.415493, 0.5234745, 0.522301, 1.... \n", - "3 0 [0.403756, 0.3826295, 0.637324, 0.5152585, 1.0... \n", - "4 0 [0.4436615, 0.3814555, 0.5938965, 0.4518785, 1... \n", - "5 0 [0.6044605, 0.3990615, 0.67253554, 0.46596247,... \n", - "6 0 [0.410798, 0.4436625, 0.556338, 0.51173747, 1.... \n", - "7 0 [0.650235, 0.4166665, 0.73826295, 0.48356748, ... \n", - "8 0 [0.64788747, 0.388498, 0.7582165, 0.49413198, ... \n", - "9 0 [0.6807515, 0.40140802, 0.75704247, 0.45774597... \n", - "10 0 [0.431925, 0.1701875, 0.561033, 0.3180745, 1.0... \n", - "11 0 [0.42018852, 0.34507048, 0.5434275, 0.4941315,... \n", - "12 0 [0.3732395, 0.348592, 0.54812247, 0.46009403, ... \n", - "13 0 [0.37793398, 0.3814555, 0.48122, 0.45892054, 1... \n", - "14 0 None \n", - " 1 [0.43427247, 0.3556335, 0.51173747, 0.43075046... \n", - " 2 [0.4929575, 0.44718304, 0.5375585, 0.485915, 1... \n", - "15 0 None \n", - " 1 [0.629108, 0.39788702, 0.67840403, 0.453051, 1... \n", - " 2 [0.4530515, 0.4107985, 0.58568054, 0.5105635, ... \n", - "16 0 None \n", - " 1 [0.45070451, 0.3967135, 0.5762915, 0.5152585, ... \n", - " 2 [0.63732404, 0.403756, 0.684272, 0.46009403, 1... \n", - "17 0 [0.43779355, 0.3826295, 0.61032856, 0.5223005,... \n", - "18 0 [0.4401405, 0.374413, 0.6068075, 0.529343, 1.0... \n", - "19 0 [0.536385, 0.4495305, 0.600939, 0.49999946, 0.... \n", - "20 0 None \n", - " 1 [0.52347445, 0.4424885, 0.6150235, 0.5281695, ... \n", - " 2 [0.456573, 0.515258, 0.504695, 0.562206, 0.0, ... \n", - " 3 [0.504695, 0.44248852, 0.54342705, 0.48004752,... \n", - "21 0 [0.54107946, 0.456573, 0.6103285, 0.511737, 0.... \n", - "22 0 [0.3298125, 0.2042255, 0.4225355, 0.2957745, 0... \n", - "23 0 [0.517606, 0.18661949, 0.615024, 0.2875585, 1.... \n", - "24 0 [0.51173747, 0.18779299, 0.6126765, 0.308685, ... \n", - "25 0 None \n", - " 1 [0.30985948, 0.3626765, 0.3861505, 0.43896753,... \n", - " 2 [0.41314548, 0.3427225, 0.48122054, 0.42018753... \n", - "26 0 [0.25821602, 0.30046952, 0.46830997, 0.4577465... \n", - "27 0 [0.2734745, 0.2828635, 0.44718352, 0.4577465, ... \n", - "28 0 [0.5915495, 0.23591551, 0.6455405, 0.2887325, ... \n", - "29 0 [0.5903755, 0.231221, 0.6420185, 0.280517, 1.0... \n", - "30 0 [0.40493003, 0.2018775, 0.469484, 0.2699525, 1... \n", - "31 0 [0.39554, 0.19483551, 0.483568, 0.2887325, 1.0... \n", - "32 0 [0.4025825, 0.212441, 0.45539945, 0.269953, 1.... \n", - "33 0 [0.2265255, 0.2136145, 0.3438965, 0.30751148, ... \n", - "34 0 [0.634976, 0.349765, 0.7723, 0.456573, 1.0, 1.0] \n", - "35 0 [0.6079815, 0.320423, 0.7969485, 0.489437, 1.0... \n", - "36 0 [0.6185445, 0.30164298, 0.7758215, 0.494131, 1... \n", - "\n", - " discarded origin task_type last_seen group_id \\\n", - "sample_id annotation_id \n", - "0 0 False train_loader -1.0 0 \n", - "1 0 False train_loader -1.0 1 \n", - "2 0 False train_loader -1.0 2 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - "3 0 False train_loader -1.0 3 \n", - "4 0 False train_loader -1.0 4 \n", - "5 0 False train_loader -1.0 5 \n", - "6 0 False train_loader -1.0 6 \n", - "7 0 False train_loader -1.0 7 \n", - "8 0 False train_loader -1.0 8 \n", - "9 0 False train_loader -1.0 9 \n", - "10 0 False train_loader -1.0 10 \n", - "11 0 False train_loader -1.0 11 \n", - "12 0 False train_loader -1.0 12 \n", - "13 0 False train_loader -1.0 13 \n", - "14 0 False train_loader -1.0 14 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - "15 0 False train_loader -1.0 15 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - "16 0 False train_loader -1.0 16 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - "17 0 False train_loader -1.0 17 \n", - "18 0 False train_loader -1.0 18 \n", - "19 0 False train_loader -1.0 19 \n", - "20 0 False train_loader -1.0 20 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - " 3 None NaN None NaN None \n", - "21 0 False train_loader -1.0 21 \n", - "22 0 False train_loader -1.0 22 \n", - "23 0 False train_loader -1.0 23 \n", - "24 0 False train_loader -1.0 24 \n", - "25 0 False train_loader -1.0 25 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - "26 0 False train_loader -1.0 26 \n", - "27 0 False train_loader -1.0 27 \n", - "28 0 False train_loader -1.0 28 \n", - "29 0 False train_loader -1.0 29 \n", - "30 0 False train_loader -1.0 30 \n", - "31 0 False train_loader -1.0 31 \n", - "32 0 False train_loader -1.0 32 \n", - "33 0 False train_loader -1.0 33 \n", - "34 0 False train_loader -1.0 34 \n", - "35 0 False train_loader -1.0 35 \n", - "36 0 False train_loader -1.0 36 \n", - "\n", - " member_rank \\\n", - "sample_id annotation_id \n", - "0 0 0.0 \n", - "1 0 0.0 \n", - "2 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - "3 0 0.0 \n", - "4 0 0.0 \n", - "5 0 0.0 \n", - "6 0 0.0 \n", - "7 0 0.0 \n", - "8 0 0.0 \n", - "9 0 0.0 \n", - "10 0 0.0 \n", - "11 0 0.0 \n", - "12 0 0.0 \n", - "13 0 0.0 \n", - "14 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - "15 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - "16 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - "17 0 0.0 \n", - "18 0 0.0 \n", - "19 0 0.0 \n", - "20 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - " 3 NaN \n", - "21 0 0.0 \n", - "22 0 0.0 \n", - "23 0 0.0 \n", - "24 0 0.0 \n", - "25 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - "26 0 0.0 \n", - "27 0 0.0 \n", - "28 0 0.0 \n", - "29 0 0.0 \n", - "30 0 0.0 \n", - "31 0 0.0 \n", - "32 0 0.0 \n", - "33 0 0.0 \n", - "34 0 0.0 \n", - "35 0 0.0 \n", - "36 0 0.0 \n", - "\n", - " img_path \\\n", - "sample_id annotation_id \n", - "0 0 /content/datasets/brain-tumor/images/train/000... \n", - "1 0 /content/datasets/brain-tumor/images/train/000... \n", - "2 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - "3 0 /content/datasets/brain-tumor/images/train/000... \n", - "4 0 /content/datasets/brain-tumor/images/train/000... \n", - "5 0 /content/datasets/brain-tumor/images/train/000... \n", - "6 0 /content/datasets/brain-tumor/images/train/000... \n", - "7 0 /content/datasets/brain-tumor/images/train/000... \n", - "8 0 /content/datasets/brain-tumor/images/train/000... \n", - "9 0 /content/datasets/brain-tumor/images/train/000... \n", - "10 0 /content/datasets/brain-tumor/images/train/000... \n", - "11 0 /content/datasets/brain-tumor/images/train/000... \n", - "12 0 /content/datasets/brain-tumor/images/train/000... \n", - "13 0 /content/datasets/brain-tumor/images/train/000... \n", - "14 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - "15 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - "16 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - "17 0 /content/datasets/brain-tumor/images/train/000... \n", - "18 0 /content/datasets/brain-tumor/images/train/000... \n", - "19 0 /content/datasets/brain-tumor/images/train/000... \n", - "20 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - " 3 None \n", - "21 0 /content/datasets/brain-tumor/images/train/000... \n", - "22 0 /content/datasets/brain-tumor/images/train/000... \n", - "23 0 /content/datasets/brain-tumor/images/train/000... \n", - "24 0 /content/datasets/brain-tumor/images/train/000... \n", - "25 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - "26 0 /content/datasets/brain-tumor/images/train/000... \n", - "27 0 /content/datasets/brain-tumor/images/train/000... \n", - "28 0 /content/datasets/brain-tumor/images/train/000... \n", - "29 0 /content/datasets/brain-tumor/images/train/000... \n", - "30 0 /content/datasets/brain-tumor/images/train/000... \n", - "31 0 /content/datasets/brain-tumor/images/train/000... \n", - "32 0 /content/datasets/brain-tumor/images/train/000... \n", - "33 0 /content/datasets/brain-tumor/images/train/000... \n", - "34 0 /content/datasets/brain-tumor/images/train/000... \n", - "35 0 /content/datasets/brain-tumor/images/train/000... \n", - "36 0 /content/datasets/brain-tumor/images/train/000... \n", - "\n", - " cls \n", - "sample_id annotation_id \n", - "0 0 [[1.0]] \n", - "1 0 [[1.0]] \n", - "2 0 [[1.0], [1.0]] \n", - " 1 None \n", - " 2 None \n", - "3 0 [[1.0]] \n", - "4 0 [[1.0]] \n", - "5 0 [[1.0]] \n", - "6 0 [[1.0]] \n", - "7 0 [[1.0]] \n", - "8 0 [[1.0]] \n", - "9 0 [[1.0]] \n", - "10 0 [[1.0]] \n", - "11 0 [[1.0]] \n", - "12 0 [[1.0]] \n", - "13 0 [[1.0]] \n", - "14 0 [[1.0], [1.0]] \n", - " 1 None \n", - " 2 None \n", - "15 0 [[1.0], [1.0]] \n", - " 1 None \n", - " 2 None \n", - "16 0 [[1.0], [1.0]] \n", - " 1 None \n", - " 2 None \n", - "17 0 [[1.0]] \n", - "18 0 [[1.0]] \n", - "19 0 [[0.0]] \n", - "20 0 [[0.0], [0.0], [0.0]] \n", - " 1 None \n", - " 2 None \n", - " 3 None \n", - "21 0 [[0.0]] \n", - "22 0 [[0.0]] \n", - "23 0 [[1.0]] \n", - "24 0 [[1.0]] \n", - "25 0 [[0.0], [0.0]] \n", - " 1 None \n", - " 2 None \n", - "26 0 [[0.0]] \n", - "27 0 [[0.0]] \n", - "28 0 [[1.0]] \n", - "29 0 [[1.0]] \n", - "30 0 [[1.0]] \n", - "31 0 [[1.0]] \n", - "32 0 [[1.0]] \n", - "33 0 [[1.0]] \n", - "34 0 [[1.0]] \n", - "35 0 [[1.0]] \n", - "36 0 [[1.0]] " - ] - }, - "metadata": {}, - "output_type": "display_data" - }, - { - "data": { - "text/html": [ - "🔗 Open the full grid in Weights Studio" - ], - "text/plain": [ - "" - ] - }, - "metadata": {}, - "output_type": "display_data" - } - ], + "outputs": [], "source": [ "import pandas as pd\n", "from IPython.display import display, HTML\n", diff --git a/weightslab/examples/Notebooks/Ultralytics/wl-how-to-train-ultralytics-yolo-on-construction-ppe-detection-dataset.ipynb b/weightslab/examples/Notebooks/Ultralytics/wl-how-to-train-ultralytics-yolo-on-construction-ppe-detection-dataset.ipynb index a68d78fb..957a8a4f 100644 --- a/weightslab/examples/Notebooks/Ultralytics/wl-how-to-train-ultralytics-yolo-on-construction-ppe-detection-dataset.ipynb +++ b/weightslab/examples/Notebooks/Ultralytics/wl-how-to-train-ultralytics-yolo-on-construction-ppe-detection-dataset.ipynb @@ -44,7 +44,7 @@ }, { "cell_type": "code", - "execution_count": 6, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/" @@ -52,24 +52,14 @@ "id": "71be394c", "outputId": "c9ef4767-84f5-4bea-e395-9e11766002f2" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Requirement already satisfied: setuptools in /usr/local/lib/python3.12/dist-packages (75.2.0)\n", - "Requirement already satisfied: wheel in /usr/local/lib/python3.12/dist-packages (0.47.0)\n", - "Requirement already satisfied: packaging>=24.0 in /usr/local/lib/python3.12/dist-packages (from wheel) (26.2)\n" - ] - } - ], + "outputs": [], "source": [ "%pip install setuptools wheel" ] }, { "cell_type": "code", - "execution_count": 7, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/", @@ -78,149 +68,7 @@ "id": "jV1oi_PQN0xK", "outputId": "eda8eca1-f4b3-4ffd-eaa6-0f67ce857153" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Collecting git+https://github.com/GrayboxTech/weightslab.git\n", - " Cloning https://github.com/GrayboxTech/weightslab.git (to revision landingcollab) to /tmp/pip-req-build-z4ba7cx9\n", - " Running command git clone --filter=blob:none --quiet https://github.com/GrayboxTech/weightslab.git /tmp/pip-req-build-z4ba7cx9\n", - " Running command git checkout -b landingcollab --track origin/landingcollab\n", - " Switched to a new branch 'landingcollab'\n", - " Branch 'landingcollab' set up to track remote branch 'landingcollab' from 'origin'.\n", - " Resolved https://github.com/GrayboxTech/weightslab.git to commit 3f1cf143bc93e67df43adc71ab831fa71477e0d2\n", - " Installing build dependencies ... \u001b[?25l\u001b[?25hdone\n", - " Getting requirements to build wheel ... \u001b[?25l\u001b[?25hdone\n", - " Preparing metadata (pyproject.toml) ... \u001b[?25l\u001b[?25hdone\n", - "Requirement already satisfied: numpy<3,>=1.24 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (2.0.2)\n", - "Requirement already satisfied: pandas<3,>=2.2.2 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (2.2.2)\n", - "Requirement already satisfied: duckdb<2,>=1.1 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.3.2)\n", - "Requirement already satisfied: PyYAML<7,>=6.0.3 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (6.0.3)\n", - "Requirement already satisfied: dill<0.5,>=0.3.8 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (0.3.8)\n", - "Requirement already satisfied: zstandard<1,>=0.22 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (0.25.0)\n", - "Requirement already satisfied: h5py<4,>=3.10 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (3.16.0)\n", - "Requirement already satisfied: xxhash<4.1,>=3.4 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (3.7.1)\n", - "Requirement already satisfied: tables<4,>=3.9 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (3.10.2)\n", - "Requirement already satisfied: torch<=2.9,>=2.1 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (2.9.0)\n", - "Requirement already satisfied: torchvision<1,>=0.16 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (0.24.0)\n", - "Requirement already satisfied: torchmetrics>=1.9 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.9.0)\n", - "Requirement already satisfied: grpcio<2,>=1.80 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.81.1)\n", - "Requirement already satisfied: protobuf<8,>=5.28.1 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (5.29.6)\n", - "Requirement already satisfied: pydantic<3,>=2.7 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (2.13.4)\n", - "Requirement already satisfied: Pillow<12,>=10 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (11.3.0)\n", - "Requirement already satisfied: graphviz<1,>=0.20 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (0.21)\n", - "Requirement already satisfied: onnx<=1.20,>=1.15 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.20.0)\n", - "Requirement already satisfied: tqdm<5,>=4.66 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (4.67.3)\n", - "Requirement already satisfied: python-dotenv<2,>=1 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.2.2)\n", - "Requirement already satisfied: langchain-core<2,>=0.3 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.4.9)\n", - "Requirement already satisfied: langchain-ollama<2,>=0.2 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.1.0)\n", - "Requirement already satisfied: langchain-openai<2,>=0.2 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.3.5)\n", - "Requirement already satisfied: typing-extensions~=4.12 in /usr/local/lib/python3.12/dist-packages (from grpcio<2,>=1.80->weightslab==1.3.3.dev8) (4.15.0)\n", - "Requirement already satisfied: jsonpatch<2.0.0,>=1.33.0 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (1.33)\n", - "Requirement already satisfied: langchain-protocol>=0.0.17 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (0.0.18)\n", - "Requirement already satisfied: langsmith<1.0.0,>=0.3.45 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (0.9.1)\n", - "Requirement already satisfied: packaging>=23.2.0 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (26.2)\n", - "Requirement already satisfied: tenacity!=8.4.0,<10.0.0,>=8.1.0 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (9.1.4)\n", - "Requirement already satisfied: uuid-utils<1.0,>=0.12.0 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (0.16.2)\n", - "Requirement already satisfied: ollama<1.0.0,>=0.6.1 in /usr/local/lib/python3.12/dist-packages (from langchain-ollama<2,>=0.2->weightslab==1.3.3.dev8) (0.6.2)\n", - "Requirement already satisfied: openai<3.0.0,>=2.45.0 in /usr/local/lib/python3.12/dist-packages (from langchain-openai<2,>=0.2->weightslab==1.3.3.dev8) (2.45.0)\n", - "Requirement already satisfied: tiktoken<1.0.0,>=0.7.0 in /usr/local/lib/python3.12/dist-packages (from langchain-openai<2,>=0.2->weightslab==1.3.3.dev8) (0.13.0)\n", - "Requirement already satisfied: ml_dtypes>=0.5.0 in /usr/local/lib/python3.12/dist-packages (from onnx<=1.20,>=1.15->weightslab==1.3.3.dev8) (0.5.4)\n", - "Requirement already satisfied: python-dateutil>=2.8.2 in /usr/local/lib/python3.12/dist-packages (from pandas<3,>=2.2.2->weightslab==1.3.3.dev8) (2.9.0.post0)\n", - "Requirement already satisfied: pytz>=2020.1 in /usr/local/lib/python3.12/dist-packages (from pandas<3,>=2.2.2->weightslab==1.3.3.dev8) (2025.2)\n", - "Requirement already satisfied: tzdata>=2022.7 in /usr/local/lib/python3.12/dist-packages (from pandas<3,>=2.2.2->weightslab==1.3.3.dev8) (2026.2)\n", - "Requirement already satisfied: annotated-types>=0.6.0 in /usr/local/lib/python3.12/dist-packages (from pydantic<3,>=2.7->weightslab==1.3.3.dev8) (0.7.0)\n", - "Requirement already satisfied: pydantic-core==2.46.4 in /usr/local/lib/python3.12/dist-packages (from pydantic<3,>=2.7->weightslab==1.3.3.dev8) (2.46.4)\n", - "Requirement already satisfied: typing-inspection>=0.4.2 in /usr/local/lib/python3.12/dist-packages (from pydantic<3,>=2.7->weightslab==1.3.3.dev8) (0.4.2)\n", - "Requirement already satisfied: numexpr>=2.6.2 in /usr/local/lib/python3.12/dist-packages (from tables<4,>=3.9->weightslab==1.3.3.dev8) (2.14.1)\n", - "Requirement already satisfied: py-cpuinfo in /usr/local/lib/python3.12/dist-packages (from tables<4,>=3.9->weightslab==1.3.3.dev8) (9.0.0)\n", - "Requirement already satisfied: blosc2>=2.3.0 in /usr/local/lib/python3.12/dist-packages (from tables<4,>=3.9->weightslab==1.3.3.dev8) (4.5.1)\n", - "Requirement already satisfied: filelock in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.29.4)\n", - "Requirement already satisfied: setuptools in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (75.2.0)\n", - "Requirement already satisfied: sympy>=1.13.3 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (1.14.0)\n", - "Requirement already satisfied: networkx>=2.5.1 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.6.1)\n", - "Requirement already satisfied: jinja2 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.1.6)\n", - "Requirement already satisfied: fsspec>=0.8.5 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (2025.3.0)\n", - "Requirement already satisfied: nvidia-cuda-nvrtc-cu12==12.8.93 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.93)\n", - "Requirement already satisfied: nvidia-cuda-runtime-cu12==12.8.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.90)\n", - "Requirement already satisfied: nvidia-cuda-cupti-cu12==12.8.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.90)\n", - "Requirement already satisfied: nvidia-cudnn-cu12==9.10.2.21 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (9.10.2.21)\n", - "Requirement already satisfied: nvidia-cublas-cu12==12.8.4.1 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.4.1)\n", - "Requirement already satisfied: nvidia-cufft-cu12==11.3.3.83 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (11.3.3.83)\n", - "Requirement already satisfied: nvidia-curand-cu12==10.3.9.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (10.3.9.90)\n", - "Requirement already satisfied: nvidia-cusolver-cu12==11.7.3.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (11.7.3.90)\n", - "Requirement already satisfied: nvidia-cusparse-cu12==12.5.8.93 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.5.8.93)\n", - "Requirement already satisfied: nvidia-cusparselt-cu12==0.7.1 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (0.7.1)\n", - "Requirement already satisfied: nvidia-nccl-cu12==2.27.5 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (2.27.5)\n", - "Requirement already satisfied: nvidia-nvshmem-cu12==3.3.20 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.3.20)\n", - "Requirement already satisfied: nvidia-nvtx-cu12==12.8.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.90)\n", - "Requirement already satisfied: nvidia-nvjitlink-cu12==12.8.93 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.93)\n", - "Requirement already satisfied: nvidia-cufile-cu12==1.13.1.3 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (1.13.1.3)\n", - "Requirement already satisfied: triton==3.5.0 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.5.0)\n", - "Requirement already satisfied: lightning-utilities>=0.15.3 in /usr/local/lib/python3.12/dist-packages (from torchmetrics>=1.9->weightslab==1.3.3.dev8) (0.15.3)\n", - "Requirement already satisfied: ndindex in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (1.10.1)\n", - "Requirement already satisfied: msgpack in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (1.2.1)\n", - "Requirement already satisfied: requests in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (2.32.4)\n", - "Requirement already satisfied: rich in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (13.9.4)\n", - "Requirement already satisfied: textual in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (6.2.1)\n", - "Requirement already satisfied: textual-plotext in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (1.0.1)\n", - "Requirement already satisfied: threadpoolctl in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (3.6.0)\n", - "Requirement already satisfied: jsonpointer>=1.9 in /usr/local/lib/python3.12/dist-packages (from jsonpatch<2.0.0,>=1.33.0->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (3.1.1)\n", - "Requirement already satisfied: anyio>=3.5.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (4.14.0)\n", - "Requirement already satisfied: distro>=1.7.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (1.9.0)\n", - "Requirement already satisfied: httpx<1,>=0.23.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (0.28.1)\n", - "Requirement already satisfied: orjson>=3.9.14 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (3.11.9)\n", - "Requirement already satisfied: requests-toolbelt>=1.0.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (1.0.0)\n", - "Requirement already satisfied: sniffio>=1.1 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (1.3.1)\n", - "Requirement already satisfied: websockets>=15.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (15.0.1)\n", - "Requirement already satisfied: jiter<1,>=0.10.0 in /usr/local/lib/python3.12/dist-packages (from openai<3.0.0,>=2.45.0->langchain-openai<2,>=0.2->weightslab==1.3.3.dev8) (0.15.0)\n", - "Requirement already satisfied: six>=1.5 in /usr/local/lib/python3.12/dist-packages (from python-dateutil>=2.8.2->pandas<3,>=2.2.2->weightslab==1.3.3.dev8) (1.17.0)\n", - "Requirement already satisfied: mpmath<1.4,>=1.1.0 in /usr/local/lib/python3.12/dist-packages (from sympy>=1.13.3->torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (1.3.0)\n", - "Requirement already satisfied: regex in /usr/local/lib/python3.12/dist-packages (from tiktoken<1.0.0,>=0.7.0->langchain-openai<2,>=0.2->weightslab==1.3.3.dev8) (2025.11.3)\n", - "Requirement already satisfied: MarkupSafe>=2.0 in /usr/local/lib/python3.12/dist-packages (from jinja2->torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.0.3)\n", - "Requirement already satisfied: idna>=2.8 in /usr/local/lib/python3.12/dist-packages (from anyio>=3.5.0->langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (3.18)\n", - "Requirement already satisfied: certifi in /usr/local/lib/python3.12/dist-packages (from httpx<1,>=0.23.0->langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (2026.6.17)\n", - "Requirement already satisfied: httpcore==1.* in /usr/local/lib/python3.12/dist-packages (from httpx<1,>=0.23.0->langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (1.0.9)\n", - "Requirement already satisfied: h11>=0.16 in /usr/local/lib/python3.12/dist-packages (from httpcore==1.*->httpx<1,>=0.23.0->langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (0.16.0)\n", - "Requirement already satisfied: charset_normalizer<4,>=2 in /usr/local/lib/python3.12/dist-packages (from requests->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (3.4.7)\n", - "Requirement already satisfied: urllib3<3,>=1.21.1 in /usr/local/lib/python3.12/dist-packages (from requests->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (2.5.0)\n", - "Requirement already satisfied: markdown-it-py>=2.2.0 in /usr/local/lib/python3.12/dist-packages (from rich->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (4.2.0)\n", - "Requirement already satisfied: pygments<3.0.0,>=2.13.0 in /usr/local/lib/python3.12/dist-packages (from rich->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (2.20.0)\n", - "Requirement already satisfied: platformdirs<5,>=3.6.0 in /usr/local/lib/python3.12/dist-packages (from textual->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (4.10.0)\n", - "Requirement already satisfied: plotext<6.0.0,>=5.2.8 in /usr/local/lib/python3.12/dist-packages (from textual-plotext->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (5.3.2)\n", - "Requirement already satisfied: mdurl~=0.1 in /usr/local/lib/python3.12/dist-packages (from markdown-it-py>=2.2.0->rich->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (0.1.2)\n", - "Requirement already satisfied: linkify-it-py<3,>=1 in /usr/local/lib/python3.12/dist-packages (from markdown-it-py[linkify,plugins]>=2.1.0->textual->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (2.1.0)\n", - "Requirement already satisfied: mdit-py-plugins>=0.5.0 in /usr/local/lib/python3.12/dist-packages (from markdown-it-py[linkify,plugins]>=2.1.0->textual->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (0.6.1)\n", - "Requirement already satisfied: uc-micro-py in /usr/local/lib/python3.12/dist-packages (from linkify-it-py<3,>=1->markdown-it-py[linkify,plugins]>=2.1.0->textual->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (2.0.0)\n", - "Building wheels for collected packages: weightslab\n", - " Building wheel for weightslab (pyproject.toml) ... \u001b[?25l\u001b[?25hdone\n", - " Created wheel for weightslab: filename=weightslab-1.3.3.dev8-py3-none-any.whl size=2827715 sha256=5c4e7c14bbf95a48b9a8c86b5a8c744af67e1eb395197cbe79d3e443c21cf86e\n", - " Stored in directory: /tmp/pip-ephem-wheel-cache-__zwgf7x/wheels/4d/c7/cd/817c2745790555e7b6c532f5b6055d47567f9416fdfc65879b\n", - "Successfully built weightslab\n", - "Installing collected packages: weightslab\n", - " Attempting uninstall: weightslab\n", - " Found existing installation: weightslab 1.3.3.dev7\n", - " Uninstalling weightslab-1.3.3.dev7:\n", - " Successfully uninstalled weightslab-1.3.3.dev7\n", - "Successfully installed weightslab-1.3.3.dev8\n" - ] - }, - { - "data": { - "application/vnd.colab-display-data+json": { - "id": "b7fce274042740d19087972fe9c01c9b", - "pip_warning": { - "packages": [ - "weightslab" - ] - } - } - }, - "metadata": {}, - "output_type": "display_data" - } - ], + "outputs": [], "source": [ "%pip install weightslab\n", "!pip install --upgrade \"protobuf>=6.31.1\" # Google Collab compat update" @@ -228,7 +76,7 @@ }, { "cell_type": "code", - "execution_count": 1, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/" @@ -236,39 +84,7 @@ "id": "pJhvREbaKs7f", "outputId": "651baa36-9ebb-49da-830b-e6c8a6e2b64f" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Ultralytics 8.4.96 🚀 Python-3.12.13 torch-2.9.0+cu128 CPU (Intel Xeon CPU @ 2.20GHz)\n", - "Setup complete ✅ (2 CPUs, 12.7 GB RAM, 30.5/107.7 GB disk)\n", - "16/07/2026-09:29:29.683 INFO:root:setup_logging: WeightsLab logging initialized - Log file: /tmp/tmp6meuzqdz/weightslab_logs/weightslab_20260716_092929.log\n", - "16/07/2026-09:29:29.685 INFO:weightslab:: WeightsLab package initialized - Log level: INFO, Log to file: True\n", - "16/07/2026-09:29:29.686 INFO:weightslab:: \n", - "╭─────────────────────────────────────────────────────────────────────────────────────────────────────╮\n", - "│ │\n", - "│ \u001b[32m$$\\ $$\\ \u001b[0m $$\\ $$\\ $$\\ \u001b[31m$$\\ \u001b[0m $$\\ │\n", - "│ \u001b[32m$$ | $\\ $$ |\u001b[0m \\__| $$ | $$ | \u001b[31m$$ | \u001b[0m $$ | │\n", - "│ \u001b[32m$$ |$$$\\ $$ |\u001b[0m $$$$$$\\ $$\\ $$$$$$\\ $$$$$$$\\ $$$$$$\\ $$$$$$$\\ \u001b[31m$$ | \u001b[0m $$$$$$\\ $$$$$$$\\ │\n", - "│ \u001b[32m$$ $$ $$\\$$ |\u001b[0m$$ __$$\\ $$ |$$ __$$\\ $$ __$$\\ \\_$$ _| $$ _____|\u001b[31m$$ | \u001b[0m \\____$$\\ $$ __$$\\ │\n", - "│ \u001b[32m$$$$ _$$$$ |\u001b[0m$$$$$$$$ |$$ |$$ / $$ |$$ | $$ | $$ | \\$$$$$$\\ \u001b[31m$$ | \u001b[0m $$$$$$$ |$$ | $$ | │\n", - "│ \u001b[32m$$$ / \\$$$ |\u001b[0m$$ ____|$$ |$$ | $$ |$$ | $$ | $$ |$$\\ \\____$$\\ \u001b[31m$$ | \u001b[0m$$ __$$ |$$ | $$ | │\n", - "│ \u001b[32m$$ / \\$$ |\u001b[0m\\$$$$$$$\\ $$ |\\$$$$$$$ |$$ | $$ | \\$$$$ |$$$$$$$ |\u001b[31m$$$$$$$$\\ \u001b[0m\\$$$$$$$ |$$$$$$$ | │\n", - "│ \u001b[32m\\__/ \\__|\u001b[0m \\_______|\\__| \\____$$ |\\__| \\__| \\____/ \\_______/ \u001b[31m\\________|\u001b[0m \\_______|\\_______/ │\n", - "│ \u001b[32m \u001b[0m $$\\ $$ |\u001b[31m\u001b[0m │\n", - "│ \u001b[32m \u001b[0m \\$$$$$$ |\u001b[31m\u001b[0m │\n", - "│ \u001b[32m \u001b[0m \\______/\u001b[31m\u001b[0m │\n", - "│ │\n", - "│ Inspect - Edit - Evolve Neural Networks │\n", - "│ │\n", - "╰─── By GrayBx, v1.3.3.dev8 ──────────────────────────────────────────────────────────────────────────╯\n", - "\n", - "\n", - "16/07/2026-09:29:29.690 INFO:weightslab.utils.telemetry:ping_import: WeightsLab uses anonymous usage data (package version used, and OS name). Set WL_NO_TELEMETRY=1 to disable.\n" - ] - } - ], + "outputs": [], "source": [ "# Install WeightsLab. Colab already ships torch, torchvision, numpy, scikit-learn\n", "# and Pillow, so nothing extra is needed here.\n", @@ -361,7 +177,6 @@ "import warnings\n", "warnings.filterwarnings(\"ignore\")\n", "\n", - "from weightslab.integrations.ultralytics import WLAwareTrainer\n", "from ultralytics import YOLO\n", "\n", "# Hyperparameters registered with WeightsLab -> live-editable from Weights Studio.\n", @@ -386,7 +201,7 @@ "# Read back the (now live) hyperparameters and hand training to WLAwareTrainer.\n", "model = YOLO(cfg[\"model\"])\n", "results = model.train(\n", - " trainer=WLAwareTrainer,\n", + " trainer=wl.WLAwareTrainer,\n", " data=str(cfg[\"data\"]),\n", " imgsz=cfg[\"image_size\"],\n", " epochs=int(cfg[\"epochs\"]),\n", @@ -416,7 +231,7 @@ }, { "cell_type": "code", - "execution_count": 3, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/", @@ -426,1230 +241,7 @@ "id": "uOJdXGA3Ks7o", "outputId": "26652fa3-6a15-43ce-9d75-034eaba21ccd" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "1241 samples tracked\n" - ] - }, - { - "data": { - "application/vnd.google.colaboratory.intrinsic+json": { - "repr_error": "Out of range float values are not JSON compliant: nan", - "type": "dataframe" - }, - "text/html": [ - "\n", - "
\n", - "
\n", - "\n", - "\n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - "
predictionprediction_rawtargetdiscardedorigintask_typelast_seengroup_idmember_rankimg_pathcls
sample_idannotation_id
00NoneNone[0.2335685, 0.254695, 0.4553995, 0.43075103, 1...Falsetrain_loader-1.000.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
10NoneNone[0.251174, 0.24061048, 0.44366202, 0.4307515, ...Falsetrain_loader-1.010.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
20NoneNoneNoneFalsetrain_loader-1.020.0/content/datasets/brain-tumor/images/train/000...[[1.0], [1.0]]
1NoneNone[0.523474, 0.3978875, 0.634976, 0.48122054, 1....NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.40610352, 0.415493, 0.5234745, 0.522301, 1....NoneNaNNoneNaNNoneNaNNoneNone
30NoneNone[0.403756, 0.3826295, 0.637324, 0.5152585, 1.0...Falsetrain_loader-1.030.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
40NoneNone[0.4436615, 0.3814555, 0.5938965, 0.4518785, 1...Falsetrain_loader-1.040.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
50NoneNone[0.6044605, 0.3990615, 0.67253554, 0.46596247,...Falsetrain_loader-1.050.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
60NoneNone[0.410798, 0.4436625, 0.556338, 0.51173747, 1....Falsetrain_loader-1.060.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
70NoneNone[0.650235, 0.4166665, 0.73826295, 0.48356748, ...Falsetrain_loader-1.070.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
80NoneNone[0.64788747, 0.388498, 0.7582165, 0.49413198, ...Falsetrain_loader-1.080.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
90NoneNone[0.6807515, 0.40140802, 0.75704247, 0.45774597...Falsetrain_loader-1.090.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
100NoneNone[0.431925, 0.1701875, 0.561033, 0.3180745, 1.0...Falsetrain_loader-1.0100.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
110NoneNone[0.42018852, 0.34507048, 0.5434275, 0.4941315,...Falsetrain_loader-1.0110.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
120NoneNone[0.3732395, 0.348592, 0.54812247, 0.46009403, ...Falsetrain_loader-1.0120.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
130NoneNone[0.37793398, 0.3814555, 0.48122, 0.45892054, 1...Falsetrain_loader-1.0130.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
140NoneNoneNoneFalsetrain_loader-1.0140.0/content/datasets/brain-tumor/images/train/000...[[1.0], [1.0]]
1NoneNone[0.43427247, 0.3556335, 0.51173747, 0.43075046...NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.4929575, 0.44718304, 0.5375585, 0.485915, 1...NoneNaNNoneNaNNoneNaNNoneNone
150NoneNoneNoneFalsetrain_loader-1.0150.0/content/datasets/brain-tumor/images/train/000...[[1.0], [1.0]]
1NoneNone[0.629108, 0.39788702, 0.67840403, 0.453051, 1...NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.4530515, 0.4107985, 0.58568054, 0.5105635, ...NoneNaNNoneNaNNoneNaNNoneNone
160NoneNoneNoneFalsetrain_loader-1.0160.0/content/datasets/brain-tumor/images/train/000...[[1.0], [1.0]]
1NoneNone[0.45070451, 0.3967135, 0.5762915, 0.5152585, ...NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.63732404, 0.403756, 0.684272, 0.46009403, 1...NoneNaNNoneNaNNoneNaNNoneNone
170NoneNone[0.43779355, 0.3826295, 0.61032856, 0.5223005,...Falsetrain_loader-1.0170.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
180NoneNone[0.4401405, 0.374413, 0.6068075, 0.529343, 1.0...Falsetrain_loader-1.0180.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
190NoneNone[0.536385, 0.4495305, 0.600939, 0.49999946, 0....Falsetrain_loader-1.0190.0/content/datasets/brain-tumor/images/train/000...[[0.0]]
200NoneNoneNoneFalsetrain_loader-1.0200.0/content/datasets/brain-tumor/images/train/000...[[0.0], [0.0], [0.0]]
1NoneNone[0.52347445, 0.4424885, 0.6150235, 0.5281695, ...NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.456573, 0.515258, 0.504695, 0.562206, 0.0, ...NoneNaNNoneNaNNoneNaNNoneNone
3NoneNone[0.504695, 0.44248852, 0.54342705, 0.48004752,...NoneNaNNoneNaNNoneNaNNoneNone
210NoneNone[0.54107946, 0.456573, 0.6103285, 0.511737, 0....Falsetrain_loader-1.0210.0/content/datasets/brain-tumor/images/train/000...[[0.0]]
220NoneNone[0.3298125, 0.2042255, 0.4225355, 0.2957745, 0...Falsetrain_loader-1.0220.0/content/datasets/brain-tumor/images/train/000...[[0.0]]
230NoneNone[0.517606, 0.18661949, 0.615024, 0.2875585, 1....Falsetrain_loader-1.0230.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
240NoneNone[0.51173747, 0.18779299, 0.6126765, 0.308685, ...Falsetrain_loader-1.0240.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
250NoneNoneNoneFalsetrain_loader-1.0250.0/content/datasets/brain-tumor/images/train/000...[[0.0], [0.0]]
1NoneNone[0.30985948, 0.3626765, 0.3861505, 0.43896753,...NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.41314548, 0.3427225, 0.48122054, 0.42018753...NoneNaNNoneNaNNoneNaNNoneNone
260NoneNone[0.25821602, 0.30046952, 0.46830997, 0.4577465...Falsetrain_loader-1.0260.0/content/datasets/brain-tumor/images/train/000...[[0.0]]
270NoneNone[0.2734745, 0.2828635, 0.44718352, 0.4577465, ...Falsetrain_loader-1.0270.0/content/datasets/brain-tumor/images/train/000...[[0.0]]
280NoneNone[0.5915495, 0.23591551, 0.6455405, 0.2887325, ...Falsetrain_loader-1.0280.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
290NoneNone[0.5903755, 0.231221, 0.6420185, 0.280517, 1.0...Falsetrain_loader-1.0290.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
300NoneNone[0.40493003, 0.2018775, 0.469484, 0.2699525, 1...Falsetrain_loader-1.0300.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
310NoneNone[0.39554, 0.19483551, 0.483568, 0.2887325, 1.0...Falsetrain_loader-1.0310.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
320NoneNone[0.4025825, 0.212441, 0.45539945, 0.269953, 1....Falsetrain_loader-1.0320.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
330NoneNone[0.2265255, 0.2136145, 0.3438965, 0.30751148, ...Falsetrain_loader-1.0330.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
340NoneNone[0.634976, 0.349765, 0.7723, 0.456573, 1.0, 1.0]Falsetrain_loader-1.0340.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
350NoneNone[0.6079815, 0.320423, 0.7969485, 0.489437, 1.0...Falsetrain_loader-1.0350.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
360NoneNone[0.6185445, 0.30164298, 0.7758215, 0.494131, 1...Falsetrain_loader-1.0360.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
\n", - "
\n", - "
\n", - "\n", - "
\n", - " \n", - "\n", - " \n", - "\n", - " \n", - "
\n", - "\n", - "\n", - "
\n", - "
\n" - ], - "text/plain": [ - " prediction prediction_raw \\\n", - "sample_id annotation_id \n", - "0 0 None None \n", - "1 0 None None \n", - "2 0 None None \n", - " 1 None None \n", - " 2 None None \n", - "3 0 None None \n", - "4 0 None None \n", - "5 0 None None \n", - "6 0 None None \n", - "7 0 None None \n", - "8 0 None None \n", - "9 0 None None \n", - "10 0 None None \n", - "11 0 None None \n", - "12 0 None None \n", - "13 0 None None \n", - "14 0 None None \n", - " 1 None None \n", - " 2 None None \n", - "15 0 None None \n", - " 1 None None \n", - " 2 None None \n", - "16 0 None None \n", - " 1 None None \n", - " 2 None None \n", - "17 0 None None \n", - "18 0 None None \n", - "19 0 None None \n", - "20 0 None None \n", - " 1 None None \n", - " 2 None None \n", - " 3 None None \n", - "21 0 None None \n", - "22 0 None None \n", - "23 0 None None \n", - "24 0 None None \n", - "25 0 None None \n", - " 1 None None \n", - " 2 None None \n", - "26 0 None None \n", - "27 0 None None \n", - "28 0 None None \n", - "29 0 None None \n", - "30 0 None None \n", - "31 0 None None \n", - "32 0 None None \n", - "33 0 None None \n", - "34 0 None None \n", - "35 0 None None \n", - "36 0 None None \n", - "\n", - " target \\\n", - "sample_id annotation_id \n", - "0 0 [0.2335685, 0.254695, 0.4553995, 0.43075103, 1... \n", - "1 0 [0.251174, 0.24061048, 0.44366202, 0.4307515, ... \n", - "2 0 None \n", - " 1 [0.523474, 0.3978875, 0.634976, 0.48122054, 1.... \n", - " 2 [0.40610352, 0.415493, 0.5234745, 0.522301, 1.... \n", - "3 0 [0.403756, 0.3826295, 0.637324, 0.5152585, 1.0... \n", - "4 0 [0.4436615, 0.3814555, 0.5938965, 0.4518785, 1... \n", - "5 0 [0.6044605, 0.3990615, 0.67253554, 0.46596247,... \n", - "6 0 [0.410798, 0.4436625, 0.556338, 0.51173747, 1.... \n", - "7 0 [0.650235, 0.4166665, 0.73826295, 0.48356748, ... \n", - "8 0 [0.64788747, 0.388498, 0.7582165, 0.49413198, ... \n", - "9 0 [0.6807515, 0.40140802, 0.75704247, 0.45774597... \n", - "10 0 [0.431925, 0.1701875, 0.561033, 0.3180745, 1.0... \n", - "11 0 [0.42018852, 0.34507048, 0.5434275, 0.4941315,... \n", - "12 0 [0.3732395, 0.348592, 0.54812247, 0.46009403, ... \n", - "13 0 [0.37793398, 0.3814555, 0.48122, 0.45892054, 1... \n", - "14 0 None \n", - " 1 [0.43427247, 0.3556335, 0.51173747, 0.43075046... \n", - " 2 [0.4929575, 0.44718304, 0.5375585, 0.485915, 1... \n", - "15 0 None \n", - " 1 [0.629108, 0.39788702, 0.67840403, 0.453051, 1... \n", - " 2 [0.4530515, 0.4107985, 0.58568054, 0.5105635, ... \n", - "16 0 None \n", - " 1 [0.45070451, 0.3967135, 0.5762915, 0.5152585, ... \n", - " 2 [0.63732404, 0.403756, 0.684272, 0.46009403, 1... \n", - "17 0 [0.43779355, 0.3826295, 0.61032856, 0.5223005,... \n", - "18 0 [0.4401405, 0.374413, 0.6068075, 0.529343, 1.0... \n", - "19 0 [0.536385, 0.4495305, 0.600939, 0.49999946, 0.... \n", - "20 0 None \n", - " 1 [0.52347445, 0.4424885, 0.6150235, 0.5281695, ... \n", - " 2 [0.456573, 0.515258, 0.504695, 0.562206, 0.0, ... \n", - " 3 [0.504695, 0.44248852, 0.54342705, 0.48004752,... \n", - "21 0 [0.54107946, 0.456573, 0.6103285, 0.511737, 0.... \n", - "22 0 [0.3298125, 0.2042255, 0.4225355, 0.2957745, 0... \n", - "23 0 [0.517606, 0.18661949, 0.615024, 0.2875585, 1.... \n", - "24 0 [0.51173747, 0.18779299, 0.6126765, 0.308685, ... \n", - "25 0 None \n", - " 1 [0.30985948, 0.3626765, 0.3861505, 0.43896753,... \n", - " 2 [0.41314548, 0.3427225, 0.48122054, 0.42018753... \n", - "26 0 [0.25821602, 0.30046952, 0.46830997, 0.4577465... \n", - "27 0 [0.2734745, 0.2828635, 0.44718352, 0.4577465, ... \n", - "28 0 [0.5915495, 0.23591551, 0.6455405, 0.2887325, ... \n", - "29 0 [0.5903755, 0.231221, 0.6420185, 0.280517, 1.0... \n", - "30 0 [0.40493003, 0.2018775, 0.469484, 0.2699525, 1... \n", - "31 0 [0.39554, 0.19483551, 0.483568, 0.2887325, 1.0... \n", - "32 0 [0.4025825, 0.212441, 0.45539945, 0.269953, 1.... \n", - "33 0 [0.2265255, 0.2136145, 0.3438965, 0.30751148, ... \n", - "34 0 [0.634976, 0.349765, 0.7723, 0.456573, 1.0, 1.0] \n", - "35 0 [0.6079815, 0.320423, 0.7969485, 0.489437, 1.0... \n", - "36 0 [0.6185445, 0.30164298, 0.7758215, 0.494131, 1... \n", - "\n", - " discarded origin task_type last_seen group_id \\\n", - "sample_id annotation_id \n", - "0 0 False train_loader -1.0 0 \n", - "1 0 False train_loader -1.0 1 \n", - "2 0 False train_loader -1.0 2 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - "3 0 False train_loader -1.0 3 \n", - "4 0 False train_loader -1.0 4 \n", - "5 0 False train_loader -1.0 5 \n", - "6 0 False train_loader -1.0 6 \n", - "7 0 False train_loader -1.0 7 \n", - "8 0 False train_loader -1.0 8 \n", - "9 0 False train_loader -1.0 9 \n", - "10 0 False train_loader -1.0 10 \n", - "11 0 False train_loader -1.0 11 \n", - "12 0 False train_loader -1.0 12 \n", - "13 0 False train_loader -1.0 13 \n", - "14 0 False train_loader -1.0 14 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - "15 0 False train_loader -1.0 15 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - "16 0 False train_loader -1.0 16 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - "17 0 False train_loader -1.0 17 \n", - "18 0 False train_loader -1.0 18 \n", - "19 0 False train_loader -1.0 19 \n", - "20 0 False train_loader -1.0 20 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - " 3 None NaN None NaN None \n", - "21 0 False train_loader -1.0 21 \n", - "22 0 False train_loader -1.0 22 \n", - "23 0 False train_loader -1.0 23 \n", - "24 0 False train_loader -1.0 24 \n", - "25 0 False train_loader -1.0 25 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - "26 0 False train_loader -1.0 26 \n", - "27 0 False train_loader -1.0 27 \n", - "28 0 False train_loader -1.0 28 \n", - "29 0 False train_loader -1.0 29 \n", - "30 0 False train_loader -1.0 30 \n", - "31 0 False train_loader -1.0 31 \n", - "32 0 False train_loader -1.0 32 \n", - "33 0 False train_loader -1.0 33 \n", - "34 0 False train_loader -1.0 34 \n", - "35 0 False train_loader -1.0 35 \n", - "36 0 False train_loader -1.0 36 \n", - "\n", - " member_rank \\\n", - "sample_id annotation_id \n", - "0 0 0.0 \n", - "1 0 0.0 \n", - "2 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - "3 0 0.0 \n", - "4 0 0.0 \n", - "5 0 0.0 \n", - "6 0 0.0 \n", - "7 0 0.0 \n", - "8 0 0.0 \n", - "9 0 0.0 \n", - "10 0 0.0 \n", - "11 0 0.0 \n", - "12 0 0.0 \n", - "13 0 0.0 \n", - "14 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - "15 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - "16 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - "17 0 0.0 \n", - "18 0 0.0 \n", - "19 0 0.0 \n", - "20 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - " 3 NaN \n", - "21 0 0.0 \n", - "22 0 0.0 \n", - "23 0 0.0 \n", - "24 0 0.0 \n", - "25 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - "26 0 0.0 \n", - "27 0 0.0 \n", - "28 0 0.0 \n", - "29 0 0.0 \n", - "30 0 0.0 \n", - "31 0 0.0 \n", - "32 0 0.0 \n", - "33 0 0.0 \n", - "34 0 0.0 \n", - "35 0 0.0 \n", - "36 0 0.0 \n", - "\n", - " img_path \\\n", - "sample_id annotation_id \n", - "0 0 /content/datasets/brain-tumor/images/train/000... \n", - "1 0 /content/datasets/brain-tumor/images/train/000... \n", - "2 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - "3 0 /content/datasets/brain-tumor/images/train/000... \n", - "4 0 /content/datasets/brain-tumor/images/train/000... \n", - "5 0 /content/datasets/brain-tumor/images/train/000... \n", - "6 0 /content/datasets/brain-tumor/images/train/000... \n", - "7 0 /content/datasets/brain-tumor/images/train/000... \n", - "8 0 /content/datasets/brain-tumor/images/train/000... \n", - "9 0 /content/datasets/brain-tumor/images/train/000... \n", - "10 0 /content/datasets/brain-tumor/images/train/000... \n", - "11 0 /content/datasets/brain-tumor/images/train/000... \n", - "12 0 /content/datasets/brain-tumor/images/train/000... \n", - "13 0 /content/datasets/brain-tumor/images/train/000... \n", - "14 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - "15 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - "16 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - "17 0 /content/datasets/brain-tumor/images/train/000... \n", - "18 0 /content/datasets/brain-tumor/images/train/000... \n", - "19 0 /content/datasets/brain-tumor/images/train/000... \n", - "20 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - " 3 None \n", - "21 0 /content/datasets/brain-tumor/images/train/000... \n", - "22 0 /content/datasets/brain-tumor/images/train/000... \n", - "23 0 /content/datasets/brain-tumor/images/train/000... \n", - "24 0 /content/datasets/brain-tumor/images/train/000... \n", - "25 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - "26 0 /content/datasets/brain-tumor/images/train/000... \n", - "27 0 /content/datasets/brain-tumor/images/train/000... \n", - "28 0 /content/datasets/brain-tumor/images/train/000... \n", - "29 0 /content/datasets/brain-tumor/images/train/000... \n", - "30 0 /content/datasets/brain-tumor/images/train/000... \n", - "31 0 /content/datasets/brain-tumor/images/train/000... \n", - "32 0 /content/datasets/brain-tumor/images/train/000... \n", - "33 0 /content/datasets/brain-tumor/images/train/000... \n", - "34 0 /content/datasets/brain-tumor/images/train/000... \n", - "35 0 /content/datasets/brain-tumor/images/train/000... \n", - "36 0 /content/datasets/brain-tumor/images/train/000... \n", - "\n", - " cls \n", - "sample_id annotation_id \n", - "0 0 [[1.0]] \n", - "1 0 [[1.0]] \n", - "2 0 [[1.0], [1.0]] \n", - " 1 None \n", - " 2 None \n", - "3 0 [[1.0]] \n", - "4 0 [[1.0]] \n", - "5 0 [[1.0]] \n", - "6 0 [[1.0]] \n", - "7 0 [[1.0]] \n", - "8 0 [[1.0]] \n", - "9 0 [[1.0]] \n", - "10 0 [[1.0]] \n", - "11 0 [[1.0]] \n", - "12 0 [[1.0]] \n", - "13 0 [[1.0]] \n", - "14 0 [[1.0], [1.0]] \n", - " 1 None \n", - " 2 None \n", - "15 0 [[1.0], [1.0]] \n", - " 1 None \n", - " 2 None \n", - "16 0 [[1.0], [1.0]] \n", - " 1 None \n", - " 2 None \n", - "17 0 [[1.0]] \n", - "18 0 [[1.0]] \n", - "19 0 [[0.0]] \n", - "20 0 [[0.0], [0.0], [0.0]] \n", - " 1 None \n", - " 2 None \n", - " 3 None \n", - "21 0 [[0.0]] \n", - "22 0 [[0.0]] \n", - "23 0 [[1.0]] \n", - "24 0 [[1.0]] \n", - "25 0 [[0.0], [0.0]] \n", - " 1 None \n", - " 2 None \n", - "26 0 [[0.0]] \n", - "27 0 [[0.0]] \n", - "28 0 [[1.0]] \n", - "29 0 [[1.0]] \n", - "30 0 [[1.0]] \n", - "31 0 [[1.0]] \n", - "32 0 [[1.0]] \n", - "33 0 [[1.0]] \n", - "34 0 [[1.0]] \n", - "35 0 [[1.0]] \n", - "36 0 [[1.0]] " - ] - }, - "metadata": {}, - "output_type": "display_data" - }, - { - "data": { - "text/html": [ - "🔗 Open the full grid in Weights Studio" - ], - "text/plain": [ - "" - ] - }, - "metadata": {}, - "output_type": "display_data" - } - ], + "outputs": [], "source": [ "import pandas as pd\n", "from IPython.display import display, HTML\n", diff --git a/weightslab/examples/Notebooks/Ultralytics/wl-how-to-train-ultralytics-yolo-on-crack-segmentation-dataset.ipynb b/weightslab/examples/Notebooks/Ultralytics/wl-how-to-train-ultralytics-yolo-on-crack-segmentation-dataset.ipynb index c20ca0cd..721ed46a 100644 --- a/weightslab/examples/Notebooks/Ultralytics/wl-how-to-train-ultralytics-yolo-on-crack-segmentation-dataset.ipynb +++ b/weightslab/examples/Notebooks/Ultralytics/wl-how-to-train-ultralytics-yolo-on-crack-segmentation-dataset.ipynb @@ -44,7 +44,7 @@ }, { "cell_type": "code", - "execution_count": 6, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/" @@ -52,24 +52,14 @@ "id": "71be394c", "outputId": "c9ef4767-84f5-4bea-e395-9e11766002f2" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Requirement already satisfied: setuptools in /usr/local/lib/python3.12/dist-packages (75.2.0)\n", - "Requirement already satisfied: wheel in /usr/local/lib/python3.12/dist-packages (0.47.0)\n", - "Requirement already satisfied: packaging>=24.0 in /usr/local/lib/python3.12/dist-packages (from wheel) (26.2)\n" - ] - } - ], + "outputs": [], "source": [ "%pip install setuptools wheel" ] }, { "cell_type": "code", - "execution_count": 7, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/", @@ -78,149 +68,7 @@ "id": "jV1oi_PQN0xK", "outputId": "eda8eca1-f4b3-4ffd-eaa6-0f67ce857153" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Collecting git+https://github.com/GrayboxTech/weightslab.git\n", - " Cloning https://github.com/GrayboxTech/weightslab.git (to revision landingcollab) to /tmp/pip-req-build-z4ba7cx9\n", - " Running command git clone --filter=blob:none --quiet https://github.com/GrayboxTech/weightslab.git /tmp/pip-req-build-z4ba7cx9\n", - " Running command git checkout -b landingcollab --track origin/landingcollab\n", - " Switched to a new branch 'landingcollab'\n", - " Branch 'landingcollab' set up to track remote branch 'landingcollab' from 'origin'.\n", - " Resolved https://github.com/GrayboxTech/weightslab.git to commit 3f1cf143bc93e67df43adc71ab831fa71477e0d2\n", - " Installing build dependencies ... \u001b[?25l\u001b[?25hdone\n", - " Getting requirements to build wheel ... \u001b[?25l\u001b[?25hdone\n", - " Preparing metadata (pyproject.toml) ... \u001b[?25l\u001b[?25hdone\n", - "Requirement already satisfied: numpy<3,>=1.24 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (2.0.2)\n", - "Requirement already satisfied: pandas<3,>=2.2.2 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (2.2.2)\n", - "Requirement already satisfied: duckdb<2,>=1.1 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.3.2)\n", - "Requirement already satisfied: PyYAML<7,>=6.0.3 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (6.0.3)\n", - "Requirement already satisfied: dill<0.5,>=0.3.8 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (0.3.8)\n", - "Requirement already satisfied: zstandard<1,>=0.22 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (0.25.0)\n", - "Requirement already satisfied: h5py<4,>=3.10 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (3.16.0)\n", - "Requirement already satisfied: xxhash<4.1,>=3.4 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (3.7.1)\n", - "Requirement already satisfied: tables<4,>=3.9 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (3.10.2)\n", - "Requirement already satisfied: torch<=2.9,>=2.1 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (2.9.0)\n", - "Requirement already satisfied: torchvision<1,>=0.16 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (0.24.0)\n", - "Requirement already satisfied: torchmetrics>=1.9 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.9.0)\n", - "Requirement already satisfied: grpcio<2,>=1.80 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.81.1)\n", - "Requirement already satisfied: protobuf<8,>=5.28.1 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (5.29.6)\n", - "Requirement already satisfied: pydantic<3,>=2.7 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (2.13.4)\n", - "Requirement already satisfied: Pillow<12,>=10 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (11.3.0)\n", - "Requirement already satisfied: graphviz<1,>=0.20 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (0.21)\n", - "Requirement already satisfied: onnx<=1.20,>=1.15 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.20.0)\n", - "Requirement already satisfied: tqdm<5,>=4.66 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (4.67.3)\n", - "Requirement already satisfied: python-dotenv<2,>=1 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.2.2)\n", - "Requirement already satisfied: langchain-core<2,>=0.3 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.4.9)\n", - "Requirement already satisfied: langchain-ollama<2,>=0.2 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.1.0)\n", - "Requirement already satisfied: langchain-openai<2,>=0.2 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.3.5)\n", - "Requirement already satisfied: typing-extensions~=4.12 in /usr/local/lib/python3.12/dist-packages (from grpcio<2,>=1.80->weightslab==1.3.3.dev8) (4.15.0)\n", - "Requirement already satisfied: jsonpatch<2.0.0,>=1.33.0 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (1.33)\n", - "Requirement already satisfied: langchain-protocol>=0.0.17 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (0.0.18)\n", - "Requirement already satisfied: langsmith<1.0.0,>=0.3.45 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (0.9.1)\n", - "Requirement already satisfied: packaging>=23.2.0 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (26.2)\n", - "Requirement already satisfied: tenacity!=8.4.0,<10.0.0,>=8.1.0 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (9.1.4)\n", - "Requirement already satisfied: uuid-utils<1.0,>=0.12.0 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (0.16.2)\n", - "Requirement already satisfied: ollama<1.0.0,>=0.6.1 in /usr/local/lib/python3.12/dist-packages (from langchain-ollama<2,>=0.2->weightslab==1.3.3.dev8) (0.6.2)\n", - "Requirement already satisfied: openai<3.0.0,>=2.45.0 in /usr/local/lib/python3.12/dist-packages (from langchain-openai<2,>=0.2->weightslab==1.3.3.dev8) (2.45.0)\n", - "Requirement already satisfied: tiktoken<1.0.0,>=0.7.0 in /usr/local/lib/python3.12/dist-packages (from langchain-openai<2,>=0.2->weightslab==1.3.3.dev8) (0.13.0)\n", - "Requirement already satisfied: ml_dtypes>=0.5.0 in /usr/local/lib/python3.12/dist-packages (from onnx<=1.20,>=1.15->weightslab==1.3.3.dev8) (0.5.4)\n", - "Requirement already satisfied: python-dateutil>=2.8.2 in /usr/local/lib/python3.12/dist-packages (from pandas<3,>=2.2.2->weightslab==1.3.3.dev8) (2.9.0.post0)\n", - "Requirement already satisfied: pytz>=2020.1 in /usr/local/lib/python3.12/dist-packages (from pandas<3,>=2.2.2->weightslab==1.3.3.dev8) (2025.2)\n", - "Requirement already satisfied: tzdata>=2022.7 in /usr/local/lib/python3.12/dist-packages (from pandas<3,>=2.2.2->weightslab==1.3.3.dev8) (2026.2)\n", - "Requirement already satisfied: annotated-types>=0.6.0 in /usr/local/lib/python3.12/dist-packages (from pydantic<3,>=2.7->weightslab==1.3.3.dev8) (0.7.0)\n", - "Requirement already satisfied: pydantic-core==2.46.4 in /usr/local/lib/python3.12/dist-packages (from pydantic<3,>=2.7->weightslab==1.3.3.dev8) (2.46.4)\n", - "Requirement already satisfied: typing-inspection>=0.4.2 in /usr/local/lib/python3.12/dist-packages (from pydantic<3,>=2.7->weightslab==1.3.3.dev8) (0.4.2)\n", - "Requirement already satisfied: numexpr>=2.6.2 in /usr/local/lib/python3.12/dist-packages (from tables<4,>=3.9->weightslab==1.3.3.dev8) (2.14.1)\n", - "Requirement already satisfied: py-cpuinfo in /usr/local/lib/python3.12/dist-packages (from tables<4,>=3.9->weightslab==1.3.3.dev8) (9.0.0)\n", - "Requirement already satisfied: blosc2>=2.3.0 in /usr/local/lib/python3.12/dist-packages (from tables<4,>=3.9->weightslab==1.3.3.dev8) (4.5.1)\n", - "Requirement already satisfied: filelock in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.29.4)\n", - "Requirement already satisfied: setuptools in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (75.2.0)\n", - "Requirement already satisfied: sympy>=1.13.3 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (1.14.0)\n", - "Requirement already satisfied: networkx>=2.5.1 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.6.1)\n", - "Requirement already satisfied: jinja2 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.1.6)\n", - "Requirement already satisfied: fsspec>=0.8.5 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (2025.3.0)\n", - "Requirement already satisfied: nvidia-cuda-nvrtc-cu12==12.8.93 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.93)\n", - "Requirement already satisfied: nvidia-cuda-runtime-cu12==12.8.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.90)\n", - "Requirement already satisfied: nvidia-cuda-cupti-cu12==12.8.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.90)\n", - "Requirement already satisfied: nvidia-cudnn-cu12==9.10.2.21 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (9.10.2.21)\n", - "Requirement already satisfied: nvidia-cublas-cu12==12.8.4.1 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.4.1)\n", - "Requirement already satisfied: nvidia-cufft-cu12==11.3.3.83 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (11.3.3.83)\n", - "Requirement already satisfied: nvidia-curand-cu12==10.3.9.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (10.3.9.90)\n", - "Requirement already satisfied: nvidia-cusolver-cu12==11.7.3.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (11.7.3.90)\n", - "Requirement already satisfied: nvidia-cusparse-cu12==12.5.8.93 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.5.8.93)\n", - "Requirement already satisfied: nvidia-cusparselt-cu12==0.7.1 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (0.7.1)\n", - "Requirement already satisfied: nvidia-nccl-cu12==2.27.5 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (2.27.5)\n", - "Requirement already satisfied: nvidia-nvshmem-cu12==3.3.20 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.3.20)\n", - "Requirement already satisfied: nvidia-nvtx-cu12==12.8.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.90)\n", - "Requirement already satisfied: nvidia-nvjitlink-cu12==12.8.93 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.93)\n", - "Requirement already satisfied: nvidia-cufile-cu12==1.13.1.3 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (1.13.1.3)\n", - "Requirement already satisfied: triton==3.5.0 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.5.0)\n", - "Requirement already satisfied: lightning-utilities>=0.15.3 in /usr/local/lib/python3.12/dist-packages (from torchmetrics>=1.9->weightslab==1.3.3.dev8) (0.15.3)\n", - "Requirement already satisfied: ndindex in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (1.10.1)\n", - "Requirement already satisfied: msgpack in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (1.2.1)\n", - "Requirement already satisfied: requests in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (2.32.4)\n", - "Requirement already satisfied: rich in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (13.9.4)\n", - "Requirement already satisfied: textual in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (6.2.1)\n", - "Requirement already satisfied: textual-plotext in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (1.0.1)\n", - "Requirement already satisfied: threadpoolctl in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (3.6.0)\n", - "Requirement already satisfied: jsonpointer>=1.9 in /usr/local/lib/python3.12/dist-packages (from jsonpatch<2.0.0,>=1.33.0->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (3.1.1)\n", - "Requirement already satisfied: anyio>=3.5.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (4.14.0)\n", - "Requirement already satisfied: distro>=1.7.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (1.9.0)\n", - "Requirement already satisfied: httpx<1,>=0.23.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (0.28.1)\n", - "Requirement already satisfied: orjson>=3.9.14 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (3.11.9)\n", - "Requirement already satisfied: requests-toolbelt>=1.0.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (1.0.0)\n", - "Requirement already satisfied: sniffio>=1.1 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (1.3.1)\n", - "Requirement already satisfied: websockets>=15.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (15.0.1)\n", - "Requirement already satisfied: jiter<1,>=0.10.0 in /usr/local/lib/python3.12/dist-packages (from openai<3.0.0,>=2.45.0->langchain-openai<2,>=0.2->weightslab==1.3.3.dev8) (0.15.0)\n", - "Requirement already satisfied: six>=1.5 in /usr/local/lib/python3.12/dist-packages (from python-dateutil>=2.8.2->pandas<3,>=2.2.2->weightslab==1.3.3.dev8) (1.17.0)\n", - "Requirement already satisfied: mpmath<1.4,>=1.1.0 in /usr/local/lib/python3.12/dist-packages (from sympy>=1.13.3->torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (1.3.0)\n", - "Requirement already satisfied: regex in /usr/local/lib/python3.12/dist-packages (from tiktoken<1.0.0,>=0.7.0->langchain-openai<2,>=0.2->weightslab==1.3.3.dev8) (2025.11.3)\n", - "Requirement already satisfied: MarkupSafe>=2.0 in /usr/local/lib/python3.12/dist-packages (from jinja2->torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.0.3)\n", - "Requirement already satisfied: idna>=2.8 in /usr/local/lib/python3.12/dist-packages (from anyio>=3.5.0->langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (3.18)\n", - "Requirement already satisfied: certifi in /usr/local/lib/python3.12/dist-packages (from httpx<1,>=0.23.0->langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (2026.6.17)\n", - "Requirement already satisfied: httpcore==1.* in /usr/local/lib/python3.12/dist-packages (from httpx<1,>=0.23.0->langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (1.0.9)\n", - "Requirement already satisfied: h11>=0.16 in /usr/local/lib/python3.12/dist-packages (from httpcore==1.*->httpx<1,>=0.23.0->langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (0.16.0)\n", - "Requirement already satisfied: charset_normalizer<4,>=2 in /usr/local/lib/python3.12/dist-packages (from requests->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (3.4.7)\n", - "Requirement already satisfied: urllib3<3,>=1.21.1 in /usr/local/lib/python3.12/dist-packages (from requests->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (2.5.0)\n", - "Requirement already satisfied: markdown-it-py>=2.2.0 in /usr/local/lib/python3.12/dist-packages (from rich->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (4.2.0)\n", - "Requirement already satisfied: pygments<3.0.0,>=2.13.0 in /usr/local/lib/python3.12/dist-packages (from rich->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (2.20.0)\n", - "Requirement already satisfied: platformdirs<5,>=3.6.0 in /usr/local/lib/python3.12/dist-packages (from textual->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (4.10.0)\n", - "Requirement already satisfied: plotext<6.0.0,>=5.2.8 in /usr/local/lib/python3.12/dist-packages (from textual-plotext->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (5.3.2)\n", - "Requirement already satisfied: mdurl~=0.1 in /usr/local/lib/python3.12/dist-packages (from markdown-it-py>=2.2.0->rich->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (0.1.2)\n", - "Requirement already satisfied: linkify-it-py<3,>=1 in /usr/local/lib/python3.12/dist-packages (from markdown-it-py[linkify,plugins]>=2.1.0->textual->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (2.1.0)\n", - "Requirement already satisfied: mdit-py-plugins>=0.5.0 in /usr/local/lib/python3.12/dist-packages (from markdown-it-py[linkify,plugins]>=2.1.0->textual->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (0.6.1)\n", - "Requirement already satisfied: uc-micro-py in /usr/local/lib/python3.12/dist-packages (from linkify-it-py<3,>=1->markdown-it-py[linkify,plugins]>=2.1.0->textual->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (2.0.0)\n", - "Building wheels for collected packages: weightslab\n", - " Building wheel for weightslab (pyproject.toml) ... \u001b[?25l\u001b[?25hdone\n", - " Created wheel for weightslab: filename=weightslab-1.3.3.dev8-py3-none-any.whl size=2827715 sha256=5c4e7c14bbf95a48b9a8c86b5a8c744af67e1eb395197cbe79d3e443c21cf86e\n", - " Stored in directory: /tmp/pip-ephem-wheel-cache-__zwgf7x/wheels/4d/c7/cd/817c2745790555e7b6c532f5b6055d47567f9416fdfc65879b\n", - "Successfully built weightslab\n", - "Installing collected packages: weightslab\n", - " Attempting uninstall: weightslab\n", - " Found existing installation: weightslab 1.3.3.dev7\n", - " Uninstalling weightslab-1.3.3.dev7:\n", - " Successfully uninstalled weightslab-1.3.3.dev7\n", - "Successfully installed weightslab-1.3.3.dev8\n" - ] - }, - { - "data": { - "application/vnd.colab-display-data+json": { - "id": "b7fce274042740d19087972fe9c01c9b", - "pip_warning": { - "packages": [ - "weightslab" - ] - } - } - }, - "metadata": {}, - "output_type": "display_data" - } - ], + "outputs": [], "source": [ "%pip install weightslab\n", "!pip install --upgrade \"protobuf>=6.31.1\" # Google Collab compat update" @@ -228,7 +76,7 @@ }, { "cell_type": "code", - "execution_count": 1, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/" @@ -236,39 +84,7 @@ "id": "pJhvREbaKs7f", "outputId": "651baa36-9ebb-49da-830b-e6c8a6e2b64f" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Ultralytics 8.4.96 🚀 Python-3.12.13 torch-2.9.0+cu128 CPU (Intel Xeon CPU @ 2.20GHz)\n", - "Setup complete ✅ (2 CPUs, 12.7 GB RAM, 30.5/107.7 GB disk)\n", - "16/07/2026-09:29:29.683 INFO:root:setup_logging: WeightsLab logging initialized - Log file: /tmp/tmp6meuzqdz/weightslab_logs/weightslab_20260716_092929.log\n", - "16/07/2026-09:29:29.685 INFO:weightslab:: WeightsLab package initialized - Log level: INFO, Log to file: True\n", - "16/07/2026-09:29:29.686 INFO:weightslab:: \n", - "╭─────────────────────────────────────────────────────────────────────────────────────────────────────╮\n", - "│ │\n", - "│ \u001b[32m$$\\ $$\\ \u001b[0m $$\\ $$\\ $$\\ \u001b[31m$$\\ \u001b[0m $$\\ │\n", - "│ \u001b[32m$$ | $\\ $$ |\u001b[0m \\__| $$ | $$ | \u001b[31m$$ | \u001b[0m $$ | │\n", - "│ \u001b[32m$$ |$$$\\ $$ |\u001b[0m $$$$$$\\ $$\\ $$$$$$\\ $$$$$$$\\ $$$$$$\\ $$$$$$$\\ \u001b[31m$$ | \u001b[0m $$$$$$\\ $$$$$$$\\ │\n", - "│ \u001b[32m$$ $$ $$\\$$ |\u001b[0m$$ __$$\\ $$ |$$ __$$\\ $$ __$$\\ \\_$$ _| $$ _____|\u001b[31m$$ | \u001b[0m \\____$$\\ $$ __$$\\ │\n", - "│ \u001b[32m$$$$ _$$$$ |\u001b[0m$$$$$$$$ |$$ |$$ / $$ |$$ | $$ | $$ | \\$$$$$$\\ \u001b[31m$$ | \u001b[0m $$$$$$$ |$$ | $$ | │\n", - "│ \u001b[32m$$$ / \\$$$ |\u001b[0m$$ ____|$$ |$$ | $$ |$$ | $$ | $$ |$$\\ \\____$$\\ \u001b[31m$$ | \u001b[0m$$ __$$ |$$ | $$ | │\n", - "│ \u001b[32m$$ / \\$$ |\u001b[0m\\$$$$$$$\\ $$ |\\$$$$$$$ |$$ | $$ | \\$$$$ |$$$$$$$ |\u001b[31m$$$$$$$$\\ \u001b[0m\\$$$$$$$ |$$$$$$$ | │\n", - "│ \u001b[32m\\__/ \\__|\u001b[0m \\_______|\\__| \\____$$ |\\__| \\__| \\____/ \\_______/ \u001b[31m\\________|\u001b[0m \\_______|\\_______/ │\n", - "│ \u001b[32m \u001b[0m $$\\ $$ |\u001b[31m\u001b[0m │\n", - "│ \u001b[32m \u001b[0m \\$$$$$$ |\u001b[31m\u001b[0m │\n", - "│ \u001b[32m \u001b[0m \\______/\u001b[31m\u001b[0m │\n", - "│ │\n", - "│ Inspect - Edit - Evolve Neural Networks │\n", - "│ │\n", - "╰─── By GrayBx, v1.3.3.dev8 ──────────────────────────────────────────────────────────────────────────╯\n", - "\n", - "\n", - "16/07/2026-09:29:29.690 INFO:weightslab.utils.telemetry:ping_import: WeightsLab uses anonymous usage data (package version used, and OS name). Set WL_NO_TELEMETRY=1 to disable.\n" - ] - } - ], + "outputs": [], "source": [ "# Install WeightsLab. Colab already ships torch, torchvision, numpy, scikit-learn\n", "# and Pillow, so nothing extra is needed here.\n", @@ -351,7 +167,6 @@ "import warnings\n", "warnings.filterwarnings(\"ignore\")\n", "\n", - "from weightslab.integrations.ultralytics import WLAwareSegmentationTrainer\n", "from ultralytics import YOLO\n", "\n", "# Hyperparameters registered with WeightsLab -> live-editable from Weights Studio.\n", @@ -376,7 +191,7 @@ "# Read back the (now live) hyperparameters and hand training to WLAwareSegmentationTrainer.\n", "model = YOLO(cfg[\"model\"])\n", "results = model.train(\n", - " trainer=WLAwareSegmentationTrainer,\n", + " trainer=wl.WLAwareSegmentationTrainer,\n", " data=str(cfg[\"data\"]),\n", " imgsz=cfg[\"image_size\"],\n", " epochs=int(cfg[\"epochs\"]),\n", @@ -406,7 +221,7 @@ }, { "cell_type": "code", - "execution_count": 3, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/", @@ -416,1230 +231,7 @@ "id": "uOJdXGA3Ks7o", "outputId": "26652fa3-6a15-43ce-9d75-034eaba21ccd" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "1241 samples tracked\n" - ] - }, - { - "data": { - "application/vnd.google.colaboratory.intrinsic+json": { - "repr_error": "Out of range float values are not JSON compliant: nan", - "type": "dataframe" - }, - "text/html": [ - "\n", - "
\n", - "
\n", - "\n", - "\n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - "
predictionprediction_rawtargetdiscardedorigintask_typelast_seengroup_idmember_rankimg_pathcls
sample_idannotation_id
00NoneNone[0.2335685, 0.254695, 0.4553995, 0.43075103, 1...Falsetrain_loader-1.000.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
10NoneNone[0.251174, 0.24061048, 0.44366202, 0.4307515, ...Falsetrain_loader-1.010.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
20NoneNoneNoneFalsetrain_loader-1.020.0/content/datasets/brain-tumor/images/train/000...[[1.0], [1.0]]
1NoneNone[0.523474, 0.3978875, 0.634976, 0.48122054, 1....NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.40610352, 0.415493, 0.5234745, 0.522301, 1....NoneNaNNoneNaNNoneNaNNoneNone
30NoneNone[0.403756, 0.3826295, 0.637324, 0.5152585, 1.0...Falsetrain_loader-1.030.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
40NoneNone[0.4436615, 0.3814555, 0.5938965, 0.4518785, 1...Falsetrain_loader-1.040.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
50NoneNone[0.6044605, 0.3990615, 0.67253554, 0.46596247,...Falsetrain_loader-1.050.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
60NoneNone[0.410798, 0.4436625, 0.556338, 0.51173747, 1....Falsetrain_loader-1.060.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
70NoneNone[0.650235, 0.4166665, 0.73826295, 0.48356748, ...Falsetrain_loader-1.070.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
80NoneNone[0.64788747, 0.388498, 0.7582165, 0.49413198, ...Falsetrain_loader-1.080.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
90NoneNone[0.6807515, 0.40140802, 0.75704247, 0.45774597...Falsetrain_loader-1.090.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
100NoneNone[0.431925, 0.1701875, 0.561033, 0.3180745, 1.0...Falsetrain_loader-1.0100.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
110NoneNone[0.42018852, 0.34507048, 0.5434275, 0.4941315,...Falsetrain_loader-1.0110.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
120NoneNone[0.3732395, 0.348592, 0.54812247, 0.46009403, ...Falsetrain_loader-1.0120.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
130NoneNone[0.37793398, 0.3814555, 0.48122, 0.45892054, 1...Falsetrain_loader-1.0130.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
140NoneNoneNoneFalsetrain_loader-1.0140.0/content/datasets/brain-tumor/images/train/000...[[1.0], [1.0]]
1NoneNone[0.43427247, 0.3556335, 0.51173747, 0.43075046...NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.4929575, 0.44718304, 0.5375585, 0.485915, 1...NoneNaNNoneNaNNoneNaNNoneNone
150NoneNoneNoneFalsetrain_loader-1.0150.0/content/datasets/brain-tumor/images/train/000...[[1.0], [1.0]]
1NoneNone[0.629108, 0.39788702, 0.67840403, 0.453051, 1...NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.4530515, 0.4107985, 0.58568054, 0.5105635, ...NoneNaNNoneNaNNoneNaNNoneNone
160NoneNoneNoneFalsetrain_loader-1.0160.0/content/datasets/brain-tumor/images/train/000...[[1.0], [1.0]]
1NoneNone[0.45070451, 0.3967135, 0.5762915, 0.5152585, ...NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.63732404, 0.403756, 0.684272, 0.46009403, 1...NoneNaNNoneNaNNoneNaNNoneNone
170NoneNone[0.43779355, 0.3826295, 0.61032856, 0.5223005,...Falsetrain_loader-1.0170.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
180NoneNone[0.4401405, 0.374413, 0.6068075, 0.529343, 1.0...Falsetrain_loader-1.0180.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
190NoneNone[0.536385, 0.4495305, 0.600939, 0.49999946, 0....Falsetrain_loader-1.0190.0/content/datasets/brain-tumor/images/train/000...[[0.0]]
200NoneNoneNoneFalsetrain_loader-1.0200.0/content/datasets/brain-tumor/images/train/000...[[0.0], [0.0], [0.0]]
1NoneNone[0.52347445, 0.4424885, 0.6150235, 0.5281695, ...NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.456573, 0.515258, 0.504695, 0.562206, 0.0, ...NoneNaNNoneNaNNoneNaNNoneNone
3NoneNone[0.504695, 0.44248852, 0.54342705, 0.48004752,...NoneNaNNoneNaNNoneNaNNoneNone
210NoneNone[0.54107946, 0.456573, 0.6103285, 0.511737, 0....Falsetrain_loader-1.0210.0/content/datasets/brain-tumor/images/train/000...[[0.0]]
220NoneNone[0.3298125, 0.2042255, 0.4225355, 0.2957745, 0...Falsetrain_loader-1.0220.0/content/datasets/brain-tumor/images/train/000...[[0.0]]
230NoneNone[0.517606, 0.18661949, 0.615024, 0.2875585, 1....Falsetrain_loader-1.0230.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
240NoneNone[0.51173747, 0.18779299, 0.6126765, 0.308685, ...Falsetrain_loader-1.0240.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
250NoneNoneNoneFalsetrain_loader-1.0250.0/content/datasets/brain-tumor/images/train/000...[[0.0], [0.0]]
1NoneNone[0.30985948, 0.3626765, 0.3861505, 0.43896753,...NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.41314548, 0.3427225, 0.48122054, 0.42018753...NoneNaNNoneNaNNoneNaNNoneNone
260NoneNone[0.25821602, 0.30046952, 0.46830997, 0.4577465...Falsetrain_loader-1.0260.0/content/datasets/brain-tumor/images/train/000...[[0.0]]
270NoneNone[0.2734745, 0.2828635, 0.44718352, 0.4577465, ...Falsetrain_loader-1.0270.0/content/datasets/brain-tumor/images/train/000...[[0.0]]
280NoneNone[0.5915495, 0.23591551, 0.6455405, 0.2887325, ...Falsetrain_loader-1.0280.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
290NoneNone[0.5903755, 0.231221, 0.6420185, 0.280517, 1.0...Falsetrain_loader-1.0290.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
300NoneNone[0.40493003, 0.2018775, 0.469484, 0.2699525, 1...Falsetrain_loader-1.0300.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
310NoneNone[0.39554, 0.19483551, 0.483568, 0.2887325, 1.0...Falsetrain_loader-1.0310.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
320NoneNone[0.4025825, 0.212441, 0.45539945, 0.269953, 1....Falsetrain_loader-1.0320.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
330NoneNone[0.2265255, 0.2136145, 0.3438965, 0.30751148, ...Falsetrain_loader-1.0330.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
340NoneNone[0.634976, 0.349765, 0.7723, 0.456573, 1.0, 1.0]Falsetrain_loader-1.0340.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
350NoneNone[0.6079815, 0.320423, 0.7969485, 0.489437, 1.0...Falsetrain_loader-1.0350.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
360NoneNone[0.6185445, 0.30164298, 0.7758215, 0.494131, 1...Falsetrain_loader-1.0360.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
\n", - "
\n", - "
\n", - "\n", - "
\n", - " \n", - "\n", - " \n", - "\n", - " \n", - "
\n", - "\n", - "\n", - "
\n", - "
\n" - ], - "text/plain": [ - " prediction prediction_raw \\\n", - "sample_id annotation_id \n", - "0 0 None None \n", - "1 0 None None \n", - "2 0 None None \n", - " 1 None None \n", - " 2 None None \n", - "3 0 None None \n", - "4 0 None None \n", - "5 0 None None \n", - "6 0 None None \n", - "7 0 None None \n", - "8 0 None None \n", - "9 0 None None \n", - "10 0 None None \n", - "11 0 None None \n", - "12 0 None None \n", - "13 0 None None \n", - "14 0 None None \n", - " 1 None None \n", - " 2 None None \n", - "15 0 None None \n", - " 1 None None \n", - " 2 None None \n", - "16 0 None None \n", - " 1 None None \n", - " 2 None None \n", - "17 0 None None \n", - "18 0 None None \n", - "19 0 None None \n", - "20 0 None None \n", - " 1 None None \n", - " 2 None None \n", - " 3 None None \n", - "21 0 None None \n", - "22 0 None None \n", - "23 0 None None \n", - "24 0 None None \n", - "25 0 None None \n", - " 1 None None \n", - " 2 None None \n", - "26 0 None None \n", - "27 0 None None \n", - "28 0 None None \n", - "29 0 None None \n", - "30 0 None None \n", - "31 0 None None \n", - "32 0 None None \n", - "33 0 None None \n", - "34 0 None None \n", - "35 0 None None \n", - "36 0 None None \n", - "\n", - " target \\\n", - "sample_id annotation_id \n", - "0 0 [0.2335685, 0.254695, 0.4553995, 0.43075103, 1... \n", - "1 0 [0.251174, 0.24061048, 0.44366202, 0.4307515, ... \n", - "2 0 None \n", - " 1 [0.523474, 0.3978875, 0.634976, 0.48122054, 1.... \n", - " 2 [0.40610352, 0.415493, 0.5234745, 0.522301, 1.... \n", - "3 0 [0.403756, 0.3826295, 0.637324, 0.5152585, 1.0... \n", - "4 0 [0.4436615, 0.3814555, 0.5938965, 0.4518785, 1... \n", - "5 0 [0.6044605, 0.3990615, 0.67253554, 0.46596247,... \n", - "6 0 [0.410798, 0.4436625, 0.556338, 0.51173747, 1.... \n", - "7 0 [0.650235, 0.4166665, 0.73826295, 0.48356748, ... \n", - "8 0 [0.64788747, 0.388498, 0.7582165, 0.49413198, ... \n", - "9 0 [0.6807515, 0.40140802, 0.75704247, 0.45774597... \n", - "10 0 [0.431925, 0.1701875, 0.561033, 0.3180745, 1.0... \n", - "11 0 [0.42018852, 0.34507048, 0.5434275, 0.4941315,... \n", - "12 0 [0.3732395, 0.348592, 0.54812247, 0.46009403, ... \n", - "13 0 [0.37793398, 0.3814555, 0.48122, 0.45892054, 1... \n", - "14 0 None \n", - " 1 [0.43427247, 0.3556335, 0.51173747, 0.43075046... \n", - " 2 [0.4929575, 0.44718304, 0.5375585, 0.485915, 1... \n", - "15 0 None \n", - " 1 [0.629108, 0.39788702, 0.67840403, 0.453051, 1... \n", - " 2 [0.4530515, 0.4107985, 0.58568054, 0.5105635, ... \n", - "16 0 None \n", - " 1 [0.45070451, 0.3967135, 0.5762915, 0.5152585, ... \n", - " 2 [0.63732404, 0.403756, 0.684272, 0.46009403, 1... \n", - "17 0 [0.43779355, 0.3826295, 0.61032856, 0.5223005,... \n", - "18 0 [0.4401405, 0.374413, 0.6068075, 0.529343, 1.0... \n", - "19 0 [0.536385, 0.4495305, 0.600939, 0.49999946, 0.... \n", - "20 0 None \n", - " 1 [0.52347445, 0.4424885, 0.6150235, 0.5281695, ... \n", - " 2 [0.456573, 0.515258, 0.504695, 0.562206, 0.0, ... \n", - " 3 [0.504695, 0.44248852, 0.54342705, 0.48004752,... \n", - "21 0 [0.54107946, 0.456573, 0.6103285, 0.511737, 0.... \n", - "22 0 [0.3298125, 0.2042255, 0.4225355, 0.2957745, 0... \n", - "23 0 [0.517606, 0.18661949, 0.615024, 0.2875585, 1.... \n", - "24 0 [0.51173747, 0.18779299, 0.6126765, 0.308685, ... \n", - "25 0 None \n", - " 1 [0.30985948, 0.3626765, 0.3861505, 0.43896753,... \n", - " 2 [0.41314548, 0.3427225, 0.48122054, 0.42018753... \n", - "26 0 [0.25821602, 0.30046952, 0.46830997, 0.4577465... \n", - "27 0 [0.2734745, 0.2828635, 0.44718352, 0.4577465, ... \n", - "28 0 [0.5915495, 0.23591551, 0.6455405, 0.2887325, ... \n", - "29 0 [0.5903755, 0.231221, 0.6420185, 0.280517, 1.0... \n", - "30 0 [0.40493003, 0.2018775, 0.469484, 0.2699525, 1... \n", - "31 0 [0.39554, 0.19483551, 0.483568, 0.2887325, 1.0... \n", - "32 0 [0.4025825, 0.212441, 0.45539945, 0.269953, 1.... \n", - "33 0 [0.2265255, 0.2136145, 0.3438965, 0.30751148, ... \n", - "34 0 [0.634976, 0.349765, 0.7723, 0.456573, 1.0, 1.0] \n", - "35 0 [0.6079815, 0.320423, 0.7969485, 0.489437, 1.0... \n", - "36 0 [0.6185445, 0.30164298, 0.7758215, 0.494131, 1... \n", - "\n", - " discarded origin task_type last_seen group_id \\\n", - "sample_id annotation_id \n", - "0 0 False train_loader -1.0 0 \n", - "1 0 False train_loader -1.0 1 \n", - "2 0 False train_loader -1.0 2 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - "3 0 False train_loader -1.0 3 \n", - "4 0 False train_loader -1.0 4 \n", - "5 0 False train_loader -1.0 5 \n", - "6 0 False train_loader -1.0 6 \n", - "7 0 False train_loader -1.0 7 \n", - "8 0 False train_loader -1.0 8 \n", - "9 0 False train_loader -1.0 9 \n", - "10 0 False train_loader -1.0 10 \n", - "11 0 False train_loader -1.0 11 \n", - "12 0 False train_loader -1.0 12 \n", - "13 0 False train_loader -1.0 13 \n", - "14 0 False train_loader -1.0 14 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - "15 0 False train_loader -1.0 15 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - "16 0 False train_loader -1.0 16 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - "17 0 False train_loader -1.0 17 \n", - "18 0 False train_loader -1.0 18 \n", - "19 0 False train_loader -1.0 19 \n", - "20 0 False train_loader -1.0 20 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - " 3 None NaN None NaN None \n", - "21 0 False train_loader -1.0 21 \n", - "22 0 False train_loader -1.0 22 \n", - "23 0 False train_loader -1.0 23 \n", - "24 0 False train_loader -1.0 24 \n", - "25 0 False train_loader -1.0 25 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - "26 0 False train_loader -1.0 26 \n", - "27 0 False train_loader -1.0 27 \n", - "28 0 False train_loader -1.0 28 \n", - "29 0 False train_loader -1.0 29 \n", - "30 0 False train_loader -1.0 30 \n", - "31 0 False train_loader -1.0 31 \n", - "32 0 False train_loader -1.0 32 \n", - "33 0 False train_loader -1.0 33 \n", - "34 0 False train_loader -1.0 34 \n", - "35 0 False train_loader -1.0 35 \n", - "36 0 False train_loader -1.0 36 \n", - "\n", - " member_rank \\\n", - "sample_id annotation_id \n", - "0 0 0.0 \n", - "1 0 0.0 \n", - "2 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - "3 0 0.0 \n", - "4 0 0.0 \n", - "5 0 0.0 \n", - "6 0 0.0 \n", - "7 0 0.0 \n", - "8 0 0.0 \n", - "9 0 0.0 \n", - "10 0 0.0 \n", - "11 0 0.0 \n", - "12 0 0.0 \n", - "13 0 0.0 \n", - "14 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - "15 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - "16 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - "17 0 0.0 \n", - "18 0 0.0 \n", - "19 0 0.0 \n", - "20 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - " 3 NaN \n", - "21 0 0.0 \n", - "22 0 0.0 \n", - "23 0 0.0 \n", - "24 0 0.0 \n", - "25 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - "26 0 0.0 \n", - "27 0 0.0 \n", - "28 0 0.0 \n", - "29 0 0.0 \n", - "30 0 0.0 \n", - "31 0 0.0 \n", - "32 0 0.0 \n", - "33 0 0.0 \n", - "34 0 0.0 \n", - "35 0 0.0 \n", - "36 0 0.0 \n", - "\n", - " img_path \\\n", - "sample_id annotation_id \n", - "0 0 /content/datasets/brain-tumor/images/train/000... \n", - "1 0 /content/datasets/brain-tumor/images/train/000... \n", - "2 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - "3 0 /content/datasets/brain-tumor/images/train/000... \n", - "4 0 /content/datasets/brain-tumor/images/train/000... \n", - "5 0 /content/datasets/brain-tumor/images/train/000... \n", - "6 0 /content/datasets/brain-tumor/images/train/000... \n", - "7 0 /content/datasets/brain-tumor/images/train/000... \n", - "8 0 /content/datasets/brain-tumor/images/train/000... \n", - "9 0 /content/datasets/brain-tumor/images/train/000... \n", - "10 0 /content/datasets/brain-tumor/images/train/000... \n", - "11 0 /content/datasets/brain-tumor/images/train/000... \n", - "12 0 /content/datasets/brain-tumor/images/train/000... \n", - "13 0 /content/datasets/brain-tumor/images/train/000... \n", - "14 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - "15 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - "16 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - "17 0 /content/datasets/brain-tumor/images/train/000... \n", - "18 0 /content/datasets/brain-tumor/images/train/000... \n", - "19 0 /content/datasets/brain-tumor/images/train/000... \n", - "20 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - " 3 None \n", - "21 0 /content/datasets/brain-tumor/images/train/000... \n", - "22 0 /content/datasets/brain-tumor/images/train/000... \n", - "23 0 /content/datasets/brain-tumor/images/train/000... \n", - "24 0 /content/datasets/brain-tumor/images/train/000... \n", - "25 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - "26 0 /content/datasets/brain-tumor/images/train/000... \n", - "27 0 /content/datasets/brain-tumor/images/train/000... \n", - "28 0 /content/datasets/brain-tumor/images/train/000... \n", - "29 0 /content/datasets/brain-tumor/images/train/000... \n", - "30 0 /content/datasets/brain-tumor/images/train/000... \n", - "31 0 /content/datasets/brain-tumor/images/train/000... \n", - "32 0 /content/datasets/brain-tumor/images/train/000... \n", - "33 0 /content/datasets/brain-tumor/images/train/000... \n", - "34 0 /content/datasets/brain-tumor/images/train/000... \n", - "35 0 /content/datasets/brain-tumor/images/train/000... \n", - "36 0 /content/datasets/brain-tumor/images/train/000... \n", - "\n", - " cls \n", - "sample_id annotation_id \n", - "0 0 [[1.0]] \n", - "1 0 [[1.0]] \n", - "2 0 [[1.0], [1.0]] \n", - " 1 None \n", - " 2 None \n", - "3 0 [[1.0]] \n", - "4 0 [[1.0]] \n", - "5 0 [[1.0]] \n", - "6 0 [[1.0]] \n", - "7 0 [[1.0]] \n", - "8 0 [[1.0]] \n", - "9 0 [[1.0]] \n", - "10 0 [[1.0]] \n", - "11 0 [[1.0]] \n", - "12 0 [[1.0]] \n", - "13 0 [[1.0]] \n", - "14 0 [[1.0], [1.0]] \n", - " 1 None \n", - " 2 None \n", - "15 0 [[1.0], [1.0]] \n", - " 1 None \n", - " 2 None \n", - "16 0 [[1.0], [1.0]] \n", - " 1 None \n", - " 2 None \n", - "17 0 [[1.0]] \n", - "18 0 [[1.0]] \n", - "19 0 [[0.0]] \n", - "20 0 [[0.0], [0.0], [0.0]] \n", - " 1 None \n", - " 2 None \n", - " 3 None \n", - "21 0 [[0.0]] \n", - "22 0 [[0.0]] \n", - "23 0 [[1.0]] \n", - "24 0 [[1.0]] \n", - "25 0 [[0.0], [0.0]] \n", - " 1 None \n", - " 2 None \n", - "26 0 [[0.0]] \n", - "27 0 [[0.0]] \n", - "28 0 [[1.0]] \n", - "29 0 [[1.0]] \n", - "30 0 [[1.0]] \n", - "31 0 [[1.0]] \n", - "32 0 [[1.0]] \n", - "33 0 [[1.0]] \n", - "34 0 [[1.0]] \n", - "35 0 [[1.0]] \n", - "36 0 [[1.0]] " - ] - }, - "metadata": {}, - "output_type": "display_data" - }, - { - "data": { - "text/html": [ - "🔗 Open the full grid in Weights Studio" - ], - "text/plain": [ - "" - ] - }, - "metadata": {}, - "output_type": "display_data" - } - ], + "outputs": [], "source": [ "import pandas as pd\n", "from IPython.display import display, HTML\n", diff --git a/weightslab/examples/Notebooks/Ultralytics/wl-how-to-train-ultralytics-yolo-on-homeobjects-dataset.ipynb b/weightslab/examples/Notebooks/Ultralytics/wl-how-to-train-ultralytics-yolo-on-homeobjects-dataset.ipynb index 88818cc6..938d5b71 100644 --- a/weightslab/examples/Notebooks/Ultralytics/wl-how-to-train-ultralytics-yolo-on-homeobjects-dataset.ipynb +++ b/weightslab/examples/Notebooks/Ultralytics/wl-how-to-train-ultralytics-yolo-on-homeobjects-dataset.ipynb @@ -44,7 +44,7 @@ }, { "cell_type": "code", - "execution_count": 6, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/" @@ -52,24 +52,14 @@ "id": "71be394c", "outputId": "c9ef4767-84f5-4bea-e395-9e11766002f2" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Requirement already satisfied: setuptools in /usr/local/lib/python3.12/dist-packages (75.2.0)\n", - "Requirement already satisfied: wheel in /usr/local/lib/python3.12/dist-packages (0.47.0)\n", - "Requirement already satisfied: packaging>=24.0 in /usr/local/lib/python3.12/dist-packages (from wheel) (26.2)\n" - ] - } - ], + "outputs": [], "source": [ "%pip install setuptools wheel" ] }, { "cell_type": "code", - "execution_count": 7, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/", @@ -78,149 +68,7 @@ "id": "jV1oi_PQN0xK", "outputId": "eda8eca1-f4b3-4ffd-eaa6-0f67ce857153" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Collecting git+https://github.com/GrayboxTech/weightslab.git\n", - " Cloning https://github.com/GrayboxTech/weightslab.git (to revision landingcollab) to /tmp/pip-req-build-z4ba7cx9\n", - " Running command git clone --filter=blob:none --quiet https://github.com/GrayboxTech/weightslab.git /tmp/pip-req-build-z4ba7cx9\n", - " Running command git checkout -b landingcollab --track origin/landingcollab\n", - " Switched to a new branch 'landingcollab'\n", - " Branch 'landingcollab' set up to track remote branch 'landingcollab' from 'origin'.\n", - " Resolved https://github.com/GrayboxTech/weightslab.git to commit 3f1cf143bc93e67df43adc71ab831fa71477e0d2\n", - " Installing build dependencies ... \u001b[?25l\u001b[?25hdone\n", - " Getting requirements to build wheel ... \u001b[?25l\u001b[?25hdone\n", - " Preparing metadata (pyproject.toml) ... \u001b[?25l\u001b[?25hdone\n", - "Requirement already satisfied: numpy<3,>=1.24 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (2.0.2)\n", - "Requirement already satisfied: pandas<3,>=2.2.2 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (2.2.2)\n", - "Requirement already satisfied: duckdb<2,>=1.1 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.3.2)\n", - "Requirement already satisfied: PyYAML<7,>=6.0.3 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (6.0.3)\n", - "Requirement already satisfied: dill<0.5,>=0.3.8 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (0.3.8)\n", - "Requirement already satisfied: zstandard<1,>=0.22 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (0.25.0)\n", - "Requirement already satisfied: h5py<4,>=3.10 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (3.16.0)\n", - "Requirement already satisfied: xxhash<4.1,>=3.4 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (3.7.1)\n", - "Requirement already satisfied: tables<4,>=3.9 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (3.10.2)\n", - "Requirement already satisfied: torch<=2.9,>=2.1 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (2.9.0)\n", - "Requirement already satisfied: torchvision<1,>=0.16 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (0.24.0)\n", - "Requirement already satisfied: torchmetrics>=1.9 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.9.0)\n", - "Requirement already satisfied: grpcio<2,>=1.80 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.81.1)\n", - "Requirement already satisfied: protobuf<8,>=5.28.1 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (5.29.6)\n", - "Requirement already satisfied: pydantic<3,>=2.7 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (2.13.4)\n", - "Requirement already satisfied: Pillow<12,>=10 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (11.3.0)\n", - "Requirement already satisfied: graphviz<1,>=0.20 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (0.21)\n", - "Requirement already satisfied: onnx<=1.20,>=1.15 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.20.0)\n", - "Requirement already satisfied: tqdm<5,>=4.66 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (4.67.3)\n", - "Requirement already satisfied: python-dotenv<2,>=1 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.2.2)\n", - "Requirement already satisfied: langchain-core<2,>=0.3 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.4.9)\n", - "Requirement already satisfied: langchain-ollama<2,>=0.2 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.1.0)\n", - "Requirement already satisfied: langchain-openai<2,>=0.2 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.3.5)\n", - "Requirement already satisfied: typing-extensions~=4.12 in /usr/local/lib/python3.12/dist-packages (from grpcio<2,>=1.80->weightslab==1.3.3.dev8) (4.15.0)\n", - "Requirement already satisfied: jsonpatch<2.0.0,>=1.33.0 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (1.33)\n", - "Requirement already satisfied: langchain-protocol>=0.0.17 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (0.0.18)\n", - "Requirement already satisfied: langsmith<1.0.0,>=0.3.45 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (0.9.1)\n", - "Requirement already satisfied: packaging>=23.2.0 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (26.2)\n", - "Requirement already satisfied: tenacity!=8.4.0,<10.0.0,>=8.1.0 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (9.1.4)\n", - "Requirement already satisfied: uuid-utils<1.0,>=0.12.0 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (0.16.2)\n", - "Requirement already satisfied: ollama<1.0.0,>=0.6.1 in /usr/local/lib/python3.12/dist-packages (from langchain-ollama<2,>=0.2->weightslab==1.3.3.dev8) (0.6.2)\n", - "Requirement already satisfied: openai<3.0.0,>=2.45.0 in /usr/local/lib/python3.12/dist-packages (from langchain-openai<2,>=0.2->weightslab==1.3.3.dev8) (2.45.0)\n", - "Requirement already satisfied: tiktoken<1.0.0,>=0.7.0 in /usr/local/lib/python3.12/dist-packages (from langchain-openai<2,>=0.2->weightslab==1.3.3.dev8) (0.13.0)\n", - "Requirement already satisfied: ml_dtypes>=0.5.0 in /usr/local/lib/python3.12/dist-packages (from onnx<=1.20,>=1.15->weightslab==1.3.3.dev8) (0.5.4)\n", - "Requirement already satisfied: python-dateutil>=2.8.2 in /usr/local/lib/python3.12/dist-packages (from pandas<3,>=2.2.2->weightslab==1.3.3.dev8) (2.9.0.post0)\n", - "Requirement already satisfied: pytz>=2020.1 in /usr/local/lib/python3.12/dist-packages (from pandas<3,>=2.2.2->weightslab==1.3.3.dev8) (2025.2)\n", - "Requirement already satisfied: tzdata>=2022.7 in /usr/local/lib/python3.12/dist-packages (from pandas<3,>=2.2.2->weightslab==1.3.3.dev8) (2026.2)\n", - "Requirement already satisfied: annotated-types>=0.6.0 in /usr/local/lib/python3.12/dist-packages (from pydantic<3,>=2.7->weightslab==1.3.3.dev8) (0.7.0)\n", - "Requirement already satisfied: pydantic-core==2.46.4 in /usr/local/lib/python3.12/dist-packages (from pydantic<3,>=2.7->weightslab==1.3.3.dev8) (2.46.4)\n", - "Requirement already satisfied: typing-inspection>=0.4.2 in /usr/local/lib/python3.12/dist-packages (from pydantic<3,>=2.7->weightslab==1.3.3.dev8) (0.4.2)\n", - "Requirement already satisfied: numexpr>=2.6.2 in /usr/local/lib/python3.12/dist-packages (from tables<4,>=3.9->weightslab==1.3.3.dev8) (2.14.1)\n", - "Requirement already satisfied: py-cpuinfo in /usr/local/lib/python3.12/dist-packages (from tables<4,>=3.9->weightslab==1.3.3.dev8) (9.0.0)\n", - "Requirement already satisfied: blosc2>=2.3.0 in /usr/local/lib/python3.12/dist-packages (from tables<4,>=3.9->weightslab==1.3.3.dev8) (4.5.1)\n", - "Requirement already satisfied: filelock in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.29.4)\n", - "Requirement already satisfied: setuptools in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (75.2.0)\n", - "Requirement already satisfied: sympy>=1.13.3 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (1.14.0)\n", - "Requirement already satisfied: networkx>=2.5.1 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.6.1)\n", - "Requirement already satisfied: jinja2 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.1.6)\n", - "Requirement already satisfied: fsspec>=0.8.5 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (2025.3.0)\n", - "Requirement already satisfied: nvidia-cuda-nvrtc-cu12==12.8.93 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.93)\n", - "Requirement already satisfied: nvidia-cuda-runtime-cu12==12.8.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.90)\n", - "Requirement already satisfied: nvidia-cuda-cupti-cu12==12.8.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.90)\n", - "Requirement already satisfied: nvidia-cudnn-cu12==9.10.2.21 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (9.10.2.21)\n", - "Requirement already satisfied: nvidia-cublas-cu12==12.8.4.1 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.4.1)\n", - "Requirement already satisfied: nvidia-cufft-cu12==11.3.3.83 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (11.3.3.83)\n", - "Requirement already satisfied: nvidia-curand-cu12==10.3.9.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (10.3.9.90)\n", - "Requirement already satisfied: nvidia-cusolver-cu12==11.7.3.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (11.7.3.90)\n", - "Requirement already satisfied: nvidia-cusparse-cu12==12.5.8.93 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.5.8.93)\n", - "Requirement already satisfied: nvidia-cusparselt-cu12==0.7.1 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (0.7.1)\n", - "Requirement already satisfied: nvidia-nccl-cu12==2.27.5 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (2.27.5)\n", - "Requirement already satisfied: nvidia-nvshmem-cu12==3.3.20 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.3.20)\n", - "Requirement already satisfied: nvidia-nvtx-cu12==12.8.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.90)\n", - "Requirement already satisfied: nvidia-nvjitlink-cu12==12.8.93 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.93)\n", - "Requirement already satisfied: nvidia-cufile-cu12==1.13.1.3 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (1.13.1.3)\n", - "Requirement already satisfied: triton==3.5.0 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.5.0)\n", - "Requirement already satisfied: lightning-utilities>=0.15.3 in /usr/local/lib/python3.12/dist-packages (from torchmetrics>=1.9->weightslab==1.3.3.dev8) (0.15.3)\n", - "Requirement already satisfied: ndindex in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (1.10.1)\n", - "Requirement already satisfied: msgpack in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (1.2.1)\n", - "Requirement already satisfied: requests in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (2.32.4)\n", - "Requirement already satisfied: rich in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (13.9.4)\n", - "Requirement already satisfied: textual in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (6.2.1)\n", - "Requirement already satisfied: textual-plotext in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (1.0.1)\n", - "Requirement already satisfied: threadpoolctl in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (3.6.0)\n", - "Requirement already satisfied: jsonpointer>=1.9 in /usr/local/lib/python3.12/dist-packages (from jsonpatch<2.0.0,>=1.33.0->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (3.1.1)\n", - "Requirement already satisfied: anyio>=3.5.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (4.14.0)\n", - "Requirement already satisfied: distro>=1.7.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (1.9.0)\n", - "Requirement already satisfied: httpx<1,>=0.23.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (0.28.1)\n", - "Requirement already satisfied: orjson>=3.9.14 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (3.11.9)\n", - "Requirement already satisfied: requests-toolbelt>=1.0.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (1.0.0)\n", - "Requirement already satisfied: sniffio>=1.1 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (1.3.1)\n", - "Requirement already satisfied: websockets>=15.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (15.0.1)\n", - "Requirement already satisfied: jiter<1,>=0.10.0 in /usr/local/lib/python3.12/dist-packages (from openai<3.0.0,>=2.45.0->langchain-openai<2,>=0.2->weightslab==1.3.3.dev8) (0.15.0)\n", - "Requirement already satisfied: six>=1.5 in /usr/local/lib/python3.12/dist-packages (from python-dateutil>=2.8.2->pandas<3,>=2.2.2->weightslab==1.3.3.dev8) (1.17.0)\n", - "Requirement already satisfied: mpmath<1.4,>=1.1.0 in /usr/local/lib/python3.12/dist-packages (from sympy>=1.13.3->torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (1.3.0)\n", - "Requirement already satisfied: regex in /usr/local/lib/python3.12/dist-packages (from tiktoken<1.0.0,>=0.7.0->langchain-openai<2,>=0.2->weightslab==1.3.3.dev8) (2025.11.3)\n", - "Requirement already satisfied: MarkupSafe>=2.0 in /usr/local/lib/python3.12/dist-packages (from jinja2->torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.0.3)\n", - "Requirement already satisfied: idna>=2.8 in /usr/local/lib/python3.12/dist-packages (from anyio>=3.5.0->langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (3.18)\n", - "Requirement already satisfied: certifi in /usr/local/lib/python3.12/dist-packages (from httpx<1,>=0.23.0->langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (2026.6.17)\n", - "Requirement already satisfied: httpcore==1.* in /usr/local/lib/python3.12/dist-packages (from httpx<1,>=0.23.0->langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (1.0.9)\n", - "Requirement already satisfied: h11>=0.16 in /usr/local/lib/python3.12/dist-packages (from httpcore==1.*->httpx<1,>=0.23.0->langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (0.16.0)\n", - "Requirement already satisfied: charset_normalizer<4,>=2 in /usr/local/lib/python3.12/dist-packages (from requests->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (3.4.7)\n", - "Requirement already satisfied: urllib3<3,>=1.21.1 in /usr/local/lib/python3.12/dist-packages (from requests->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (2.5.0)\n", - "Requirement already satisfied: markdown-it-py>=2.2.0 in /usr/local/lib/python3.12/dist-packages (from rich->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (4.2.0)\n", - "Requirement already satisfied: pygments<3.0.0,>=2.13.0 in /usr/local/lib/python3.12/dist-packages (from rich->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (2.20.0)\n", - "Requirement already satisfied: platformdirs<5,>=3.6.0 in /usr/local/lib/python3.12/dist-packages (from textual->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (4.10.0)\n", - "Requirement already satisfied: plotext<6.0.0,>=5.2.8 in /usr/local/lib/python3.12/dist-packages (from textual-plotext->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (5.3.2)\n", - "Requirement already satisfied: mdurl~=0.1 in /usr/local/lib/python3.12/dist-packages (from markdown-it-py>=2.2.0->rich->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (0.1.2)\n", - "Requirement already satisfied: linkify-it-py<3,>=1 in /usr/local/lib/python3.12/dist-packages (from markdown-it-py[linkify,plugins]>=2.1.0->textual->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (2.1.0)\n", - "Requirement already satisfied: mdit-py-plugins>=0.5.0 in /usr/local/lib/python3.12/dist-packages (from markdown-it-py[linkify,plugins]>=2.1.0->textual->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (0.6.1)\n", - "Requirement already satisfied: uc-micro-py in /usr/local/lib/python3.12/dist-packages (from linkify-it-py<3,>=1->markdown-it-py[linkify,plugins]>=2.1.0->textual->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (2.0.0)\n", - "Building wheels for collected packages: weightslab\n", - " Building wheel for weightslab (pyproject.toml) ... \u001b[?25l\u001b[?25hdone\n", - " Created wheel for weightslab: filename=weightslab-1.3.3.dev8-py3-none-any.whl size=2827715 sha256=5c4e7c14bbf95a48b9a8c86b5a8c744af67e1eb395197cbe79d3e443c21cf86e\n", - " Stored in directory: /tmp/pip-ephem-wheel-cache-__zwgf7x/wheels/4d/c7/cd/817c2745790555e7b6c532f5b6055d47567f9416fdfc65879b\n", - "Successfully built weightslab\n", - "Installing collected packages: weightslab\n", - " Attempting uninstall: weightslab\n", - " Found existing installation: weightslab 1.3.3.dev7\n", - " Uninstalling weightslab-1.3.3.dev7:\n", - " Successfully uninstalled weightslab-1.3.3.dev7\n", - "Successfully installed weightslab-1.3.3.dev8\n" - ] - }, - { - "data": { - "application/vnd.colab-display-data+json": { - "id": "b7fce274042740d19087972fe9c01c9b", - "pip_warning": { - "packages": [ - "weightslab" - ] - } - } - }, - "metadata": {}, - "output_type": "display_data" - } - ], + "outputs": [], "source": [ "%pip install weightslab\n", "!pip install --upgrade \"protobuf>=6.31.1\" # Google Collab compat update" @@ -228,7 +76,7 @@ }, { "cell_type": "code", - "execution_count": 1, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/" @@ -236,39 +84,7 @@ "id": "pJhvREbaKs7f", "outputId": "651baa36-9ebb-49da-830b-e6c8a6e2b64f" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Ultralytics 8.4.96 🚀 Python-3.12.13 torch-2.9.0+cu128 CPU (Intel Xeon CPU @ 2.20GHz)\n", - "Setup complete ✅ (2 CPUs, 12.7 GB RAM, 30.5/107.7 GB disk)\n", - "16/07/2026-09:29:29.683 INFO:root:setup_logging: WeightsLab logging initialized - Log file: /tmp/tmp6meuzqdz/weightslab_logs/weightslab_20260716_092929.log\n", - "16/07/2026-09:29:29.685 INFO:weightslab:: WeightsLab package initialized - Log level: INFO, Log to file: True\n", - "16/07/2026-09:29:29.686 INFO:weightslab:: \n", - "╭─────────────────────────────────────────────────────────────────────────────────────────────────────╮\n", - "│ │\n", - "│ \u001b[32m$$\\ $$\\ \u001b[0m $$\\ $$\\ $$\\ \u001b[31m$$\\ \u001b[0m $$\\ │\n", - "│ \u001b[32m$$ | $\\ $$ |\u001b[0m \\__| $$ | $$ | \u001b[31m$$ | \u001b[0m $$ | │\n", - "│ \u001b[32m$$ |$$$\\ $$ |\u001b[0m $$$$$$\\ $$\\ $$$$$$\\ $$$$$$$\\ $$$$$$\\ $$$$$$$\\ \u001b[31m$$ | \u001b[0m $$$$$$\\ $$$$$$$\\ │\n", - "│ \u001b[32m$$ $$ $$\\$$ |\u001b[0m$$ __$$\\ $$ |$$ __$$\\ $$ __$$\\ \\_$$ _| $$ _____|\u001b[31m$$ | \u001b[0m \\____$$\\ $$ __$$\\ │\n", - "│ \u001b[32m$$$$ _$$$$ |\u001b[0m$$$$$$$$ |$$ |$$ / $$ |$$ | $$ | $$ | \\$$$$$$\\ \u001b[31m$$ | \u001b[0m $$$$$$$ |$$ | $$ | │\n", - "│ \u001b[32m$$$ / \\$$$ |\u001b[0m$$ ____|$$ |$$ | $$ |$$ | $$ | $$ |$$\\ \\____$$\\ \u001b[31m$$ | \u001b[0m$$ __$$ |$$ | $$ | │\n", - "│ \u001b[32m$$ / \\$$ |\u001b[0m\\$$$$$$$\\ $$ |\\$$$$$$$ |$$ | $$ | \\$$$$ |$$$$$$$ |\u001b[31m$$$$$$$$\\ \u001b[0m\\$$$$$$$ |$$$$$$$ | │\n", - "│ \u001b[32m\\__/ \\__|\u001b[0m \\_______|\\__| \\____$$ |\\__| \\__| \\____/ \\_______/ \u001b[31m\\________|\u001b[0m \\_______|\\_______/ │\n", - "│ \u001b[32m \u001b[0m $$\\ $$ |\u001b[31m\u001b[0m │\n", - "│ \u001b[32m \u001b[0m \\$$$$$$ |\u001b[31m\u001b[0m │\n", - "│ \u001b[32m \u001b[0m \\______/\u001b[31m\u001b[0m │\n", - "│ │\n", - "│ Inspect - Edit - Evolve Neural Networks │\n", - "│ │\n", - "╰─── By GrayBx, v1.3.3.dev8 ──────────────────────────────────────────────────────────────────────────╯\n", - "\n", - "\n", - "16/07/2026-09:29:29.690 INFO:weightslab.utils.telemetry:ping_import: WeightsLab uses anonymous usage data (package version used, and OS name). Set WL_NO_TELEMETRY=1 to disable.\n" - ] - } - ], + "outputs": [], "source": [ "# Install WeightsLab. Colab already ships torch, torchvision, numpy, scikit-learn\n", "# and Pillow, so nothing extra is needed here.\n", @@ -361,7 +177,6 @@ "import warnings\n", "warnings.filterwarnings(\"ignore\")\n", "\n", - "from weightslab.integrations.ultralytics import WLAwareTrainer\n", "from ultralytics import YOLO\n", "\n", "# Hyperparameters registered with WeightsLab -> live-editable from Weights Studio.\n", @@ -386,7 +201,7 @@ "# Read back the (now live) hyperparameters and hand training to WLAwareTrainer.\n", "model = YOLO(cfg[\"model\"])\n", "results = model.train(\n", - " trainer=WLAwareTrainer,\n", + " trainer=wl.WLAwareTrainer,\n", " data=str(cfg[\"data\"]),\n", " imgsz=cfg[\"image_size\"],\n", " epochs=int(cfg[\"epochs\"]),\n", @@ -416,7 +231,7 @@ }, { "cell_type": "code", - "execution_count": 3, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/", @@ -426,1230 +241,7 @@ "id": "uOJdXGA3Ks7o", "outputId": "26652fa3-6a15-43ce-9d75-034eaba21ccd" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "1241 samples tracked\n" - ] - }, - { - "data": { - "application/vnd.google.colaboratory.intrinsic+json": { - "repr_error": "Out of range float values are not JSON compliant: nan", - "type": "dataframe" - }, - "text/html": [ - "\n", - "
\n", - "
\n", - "\n", - "\n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - "
predictionprediction_rawtargetdiscardedorigintask_typelast_seengroup_idmember_rankimg_pathcls
sample_idannotation_id
00NoneNone[0.2335685, 0.254695, 0.4553995, 0.43075103, 1...Falsetrain_loader-1.000.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
10NoneNone[0.251174, 0.24061048, 0.44366202, 0.4307515, ...Falsetrain_loader-1.010.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
20NoneNoneNoneFalsetrain_loader-1.020.0/content/datasets/brain-tumor/images/train/000...[[1.0], [1.0]]
1NoneNone[0.523474, 0.3978875, 0.634976, 0.48122054, 1....NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.40610352, 0.415493, 0.5234745, 0.522301, 1....NoneNaNNoneNaNNoneNaNNoneNone
30NoneNone[0.403756, 0.3826295, 0.637324, 0.5152585, 1.0...Falsetrain_loader-1.030.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
40NoneNone[0.4436615, 0.3814555, 0.5938965, 0.4518785, 1...Falsetrain_loader-1.040.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
50NoneNone[0.6044605, 0.3990615, 0.67253554, 0.46596247,...Falsetrain_loader-1.050.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
60NoneNone[0.410798, 0.4436625, 0.556338, 0.51173747, 1....Falsetrain_loader-1.060.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
70NoneNone[0.650235, 0.4166665, 0.73826295, 0.48356748, ...Falsetrain_loader-1.070.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
80NoneNone[0.64788747, 0.388498, 0.7582165, 0.49413198, ...Falsetrain_loader-1.080.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
90NoneNone[0.6807515, 0.40140802, 0.75704247, 0.45774597...Falsetrain_loader-1.090.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
100NoneNone[0.431925, 0.1701875, 0.561033, 0.3180745, 1.0...Falsetrain_loader-1.0100.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
110NoneNone[0.42018852, 0.34507048, 0.5434275, 0.4941315,...Falsetrain_loader-1.0110.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
120NoneNone[0.3732395, 0.348592, 0.54812247, 0.46009403, ...Falsetrain_loader-1.0120.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
130NoneNone[0.37793398, 0.3814555, 0.48122, 0.45892054, 1...Falsetrain_loader-1.0130.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
140NoneNoneNoneFalsetrain_loader-1.0140.0/content/datasets/brain-tumor/images/train/000...[[1.0], [1.0]]
1NoneNone[0.43427247, 0.3556335, 0.51173747, 0.43075046...NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.4929575, 0.44718304, 0.5375585, 0.485915, 1...NoneNaNNoneNaNNoneNaNNoneNone
150NoneNoneNoneFalsetrain_loader-1.0150.0/content/datasets/brain-tumor/images/train/000...[[1.0], [1.0]]
1NoneNone[0.629108, 0.39788702, 0.67840403, 0.453051, 1...NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.4530515, 0.4107985, 0.58568054, 0.5105635, ...NoneNaNNoneNaNNoneNaNNoneNone
160NoneNoneNoneFalsetrain_loader-1.0160.0/content/datasets/brain-tumor/images/train/000...[[1.0], [1.0]]
1NoneNone[0.45070451, 0.3967135, 0.5762915, 0.5152585, ...NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.63732404, 0.403756, 0.684272, 0.46009403, 1...NoneNaNNoneNaNNoneNaNNoneNone
170NoneNone[0.43779355, 0.3826295, 0.61032856, 0.5223005,...Falsetrain_loader-1.0170.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
180NoneNone[0.4401405, 0.374413, 0.6068075, 0.529343, 1.0...Falsetrain_loader-1.0180.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
190NoneNone[0.536385, 0.4495305, 0.600939, 0.49999946, 0....Falsetrain_loader-1.0190.0/content/datasets/brain-tumor/images/train/000...[[0.0]]
200NoneNoneNoneFalsetrain_loader-1.0200.0/content/datasets/brain-tumor/images/train/000...[[0.0], [0.0], [0.0]]
1NoneNone[0.52347445, 0.4424885, 0.6150235, 0.5281695, ...NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.456573, 0.515258, 0.504695, 0.562206, 0.0, ...NoneNaNNoneNaNNoneNaNNoneNone
3NoneNone[0.504695, 0.44248852, 0.54342705, 0.48004752,...NoneNaNNoneNaNNoneNaNNoneNone
210NoneNone[0.54107946, 0.456573, 0.6103285, 0.511737, 0....Falsetrain_loader-1.0210.0/content/datasets/brain-tumor/images/train/000...[[0.0]]
220NoneNone[0.3298125, 0.2042255, 0.4225355, 0.2957745, 0...Falsetrain_loader-1.0220.0/content/datasets/brain-tumor/images/train/000...[[0.0]]
230NoneNone[0.517606, 0.18661949, 0.615024, 0.2875585, 1....Falsetrain_loader-1.0230.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
240NoneNone[0.51173747, 0.18779299, 0.6126765, 0.308685, ...Falsetrain_loader-1.0240.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
250NoneNoneNoneFalsetrain_loader-1.0250.0/content/datasets/brain-tumor/images/train/000...[[0.0], [0.0]]
1NoneNone[0.30985948, 0.3626765, 0.3861505, 0.43896753,...NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.41314548, 0.3427225, 0.48122054, 0.42018753...NoneNaNNoneNaNNoneNaNNoneNone
260NoneNone[0.25821602, 0.30046952, 0.46830997, 0.4577465...Falsetrain_loader-1.0260.0/content/datasets/brain-tumor/images/train/000...[[0.0]]
270NoneNone[0.2734745, 0.2828635, 0.44718352, 0.4577465, ...Falsetrain_loader-1.0270.0/content/datasets/brain-tumor/images/train/000...[[0.0]]
280NoneNone[0.5915495, 0.23591551, 0.6455405, 0.2887325, ...Falsetrain_loader-1.0280.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
290NoneNone[0.5903755, 0.231221, 0.6420185, 0.280517, 1.0...Falsetrain_loader-1.0290.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
300NoneNone[0.40493003, 0.2018775, 0.469484, 0.2699525, 1...Falsetrain_loader-1.0300.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
310NoneNone[0.39554, 0.19483551, 0.483568, 0.2887325, 1.0...Falsetrain_loader-1.0310.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
320NoneNone[0.4025825, 0.212441, 0.45539945, 0.269953, 1....Falsetrain_loader-1.0320.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
330NoneNone[0.2265255, 0.2136145, 0.3438965, 0.30751148, ...Falsetrain_loader-1.0330.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
340NoneNone[0.634976, 0.349765, 0.7723, 0.456573, 1.0, 1.0]Falsetrain_loader-1.0340.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
350NoneNone[0.6079815, 0.320423, 0.7969485, 0.489437, 1.0...Falsetrain_loader-1.0350.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
360NoneNone[0.6185445, 0.30164298, 0.7758215, 0.494131, 1...Falsetrain_loader-1.0360.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
\n", - "
\n", - "
\n", - "\n", - "
\n", - " \n", - "\n", - " \n", - "\n", - " \n", - "
\n", - "\n", - "\n", - "
\n", - "
\n" - ], - "text/plain": [ - " prediction prediction_raw \\\n", - "sample_id annotation_id \n", - "0 0 None None \n", - "1 0 None None \n", - "2 0 None None \n", - " 1 None None \n", - " 2 None None \n", - "3 0 None None \n", - "4 0 None None \n", - "5 0 None None \n", - "6 0 None None \n", - "7 0 None None \n", - "8 0 None None \n", - "9 0 None None \n", - "10 0 None None \n", - "11 0 None None \n", - "12 0 None None \n", - "13 0 None None \n", - "14 0 None None \n", - " 1 None None \n", - " 2 None None \n", - "15 0 None None \n", - " 1 None None \n", - " 2 None None \n", - "16 0 None None \n", - " 1 None None \n", - " 2 None None \n", - "17 0 None None \n", - "18 0 None None \n", - "19 0 None None \n", - "20 0 None None \n", - " 1 None None \n", - " 2 None None \n", - " 3 None None \n", - "21 0 None None \n", - "22 0 None None \n", - "23 0 None None \n", - "24 0 None None \n", - "25 0 None None \n", - " 1 None None \n", - " 2 None None \n", - "26 0 None None \n", - "27 0 None None \n", - "28 0 None None \n", - "29 0 None None \n", - "30 0 None None \n", - "31 0 None None \n", - "32 0 None None \n", - "33 0 None None \n", - "34 0 None None \n", - "35 0 None None \n", - "36 0 None None \n", - "\n", - " target \\\n", - "sample_id annotation_id \n", - "0 0 [0.2335685, 0.254695, 0.4553995, 0.43075103, 1... \n", - "1 0 [0.251174, 0.24061048, 0.44366202, 0.4307515, ... \n", - "2 0 None \n", - " 1 [0.523474, 0.3978875, 0.634976, 0.48122054, 1.... \n", - " 2 [0.40610352, 0.415493, 0.5234745, 0.522301, 1.... \n", - "3 0 [0.403756, 0.3826295, 0.637324, 0.5152585, 1.0... \n", - "4 0 [0.4436615, 0.3814555, 0.5938965, 0.4518785, 1... \n", - "5 0 [0.6044605, 0.3990615, 0.67253554, 0.46596247,... \n", - "6 0 [0.410798, 0.4436625, 0.556338, 0.51173747, 1.... \n", - "7 0 [0.650235, 0.4166665, 0.73826295, 0.48356748, ... \n", - "8 0 [0.64788747, 0.388498, 0.7582165, 0.49413198, ... \n", - "9 0 [0.6807515, 0.40140802, 0.75704247, 0.45774597... \n", - "10 0 [0.431925, 0.1701875, 0.561033, 0.3180745, 1.0... \n", - "11 0 [0.42018852, 0.34507048, 0.5434275, 0.4941315,... \n", - "12 0 [0.3732395, 0.348592, 0.54812247, 0.46009403, ... \n", - "13 0 [0.37793398, 0.3814555, 0.48122, 0.45892054, 1... \n", - "14 0 None \n", - " 1 [0.43427247, 0.3556335, 0.51173747, 0.43075046... \n", - " 2 [0.4929575, 0.44718304, 0.5375585, 0.485915, 1... \n", - "15 0 None \n", - " 1 [0.629108, 0.39788702, 0.67840403, 0.453051, 1... \n", - " 2 [0.4530515, 0.4107985, 0.58568054, 0.5105635, ... \n", - "16 0 None \n", - " 1 [0.45070451, 0.3967135, 0.5762915, 0.5152585, ... \n", - " 2 [0.63732404, 0.403756, 0.684272, 0.46009403, 1... \n", - "17 0 [0.43779355, 0.3826295, 0.61032856, 0.5223005,... \n", - "18 0 [0.4401405, 0.374413, 0.6068075, 0.529343, 1.0... \n", - "19 0 [0.536385, 0.4495305, 0.600939, 0.49999946, 0.... \n", - "20 0 None \n", - " 1 [0.52347445, 0.4424885, 0.6150235, 0.5281695, ... \n", - " 2 [0.456573, 0.515258, 0.504695, 0.562206, 0.0, ... \n", - " 3 [0.504695, 0.44248852, 0.54342705, 0.48004752,... \n", - "21 0 [0.54107946, 0.456573, 0.6103285, 0.511737, 0.... \n", - "22 0 [0.3298125, 0.2042255, 0.4225355, 0.2957745, 0... \n", - "23 0 [0.517606, 0.18661949, 0.615024, 0.2875585, 1.... \n", - "24 0 [0.51173747, 0.18779299, 0.6126765, 0.308685, ... \n", - "25 0 None \n", - " 1 [0.30985948, 0.3626765, 0.3861505, 0.43896753,... \n", - " 2 [0.41314548, 0.3427225, 0.48122054, 0.42018753... \n", - "26 0 [0.25821602, 0.30046952, 0.46830997, 0.4577465... \n", - "27 0 [0.2734745, 0.2828635, 0.44718352, 0.4577465, ... \n", - "28 0 [0.5915495, 0.23591551, 0.6455405, 0.2887325, ... \n", - "29 0 [0.5903755, 0.231221, 0.6420185, 0.280517, 1.0... \n", - "30 0 [0.40493003, 0.2018775, 0.469484, 0.2699525, 1... \n", - "31 0 [0.39554, 0.19483551, 0.483568, 0.2887325, 1.0... \n", - "32 0 [0.4025825, 0.212441, 0.45539945, 0.269953, 1.... \n", - "33 0 [0.2265255, 0.2136145, 0.3438965, 0.30751148, ... \n", - "34 0 [0.634976, 0.349765, 0.7723, 0.456573, 1.0, 1.0] \n", - "35 0 [0.6079815, 0.320423, 0.7969485, 0.489437, 1.0... \n", - "36 0 [0.6185445, 0.30164298, 0.7758215, 0.494131, 1... \n", - "\n", - " discarded origin task_type last_seen group_id \\\n", - "sample_id annotation_id \n", - "0 0 False train_loader -1.0 0 \n", - "1 0 False train_loader -1.0 1 \n", - "2 0 False train_loader -1.0 2 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - "3 0 False train_loader -1.0 3 \n", - "4 0 False train_loader -1.0 4 \n", - "5 0 False train_loader -1.0 5 \n", - "6 0 False train_loader -1.0 6 \n", - "7 0 False train_loader -1.0 7 \n", - "8 0 False train_loader -1.0 8 \n", - "9 0 False train_loader -1.0 9 \n", - "10 0 False train_loader -1.0 10 \n", - "11 0 False train_loader -1.0 11 \n", - "12 0 False train_loader -1.0 12 \n", - "13 0 False train_loader -1.0 13 \n", - "14 0 False train_loader -1.0 14 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - "15 0 False train_loader -1.0 15 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - "16 0 False train_loader -1.0 16 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - "17 0 False train_loader -1.0 17 \n", - "18 0 False train_loader -1.0 18 \n", - "19 0 False train_loader -1.0 19 \n", - "20 0 False train_loader -1.0 20 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - " 3 None NaN None NaN None \n", - "21 0 False train_loader -1.0 21 \n", - "22 0 False train_loader -1.0 22 \n", - "23 0 False train_loader -1.0 23 \n", - "24 0 False train_loader -1.0 24 \n", - "25 0 False train_loader -1.0 25 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - "26 0 False train_loader -1.0 26 \n", - "27 0 False train_loader -1.0 27 \n", - "28 0 False train_loader -1.0 28 \n", - "29 0 False train_loader -1.0 29 \n", - "30 0 False train_loader -1.0 30 \n", - "31 0 False train_loader -1.0 31 \n", - "32 0 False train_loader -1.0 32 \n", - "33 0 False train_loader -1.0 33 \n", - "34 0 False train_loader -1.0 34 \n", - "35 0 False train_loader -1.0 35 \n", - "36 0 False train_loader -1.0 36 \n", - "\n", - " member_rank \\\n", - "sample_id annotation_id \n", - "0 0 0.0 \n", - "1 0 0.0 \n", - "2 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - "3 0 0.0 \n", - "4 0 0.0 \n", - "5 0 0.0 \n", - "6 0 0.0 \n", - "7 0 0.0 \n", - "8 0 0.0 \n", - "9 0 0.0 \n", - "10 0 0.0 \n", - "11 0 0.0 \n", - "12 0 0.0 \n", - "13 0 0.0 \n", - "14 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - "15 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - "16 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - "17 0 0.0 \n", - "18 0 0.0 \n", - "19 0 0.0 \n", - "20 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - " 3 NaN \n", - "21 0 0.0 \n", - "22 0 0.0 \n", - "23 0 0.0 \n", - "24 0 0.0 \n", - "25 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - "26 0 0.0 \n", - "27 0 0.0 \n", - "28 0 0.0 \n", - "29 0 0.0 \n", - "30 0 0.0 \n", - "31 0 0.0 \n", - "32 0 0.0 \n", - "33 0 0.0 \n", - "34 0 0.0 \n", - "35 0 0.0 \n", - "36 0 0.0 \n", - "\n", - " img_path \\\n", - "sample_id annotation_id \n", - "0 0 /content/datasets/brain-tumor/images/train/000... \n", - "1 0 /content/datasets/brain-tumor/images/train/000... \n", - "2 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - "3 0 /content/datasets/brain-tumor/images/train/000... \n", - "4 0 /content/datasets/brain-tumor/images/train/000... \n", - "5 0 /content/datasets/brain-tumor/images/train/000... \n", - "6 0 /content/datasets/brain-tumor/images/train/000... \n", - "7 0 /content/datasets/brain-tumor/images/train/000... \n", - "8 0 /content/datasets/brain-tumor/images/train/000... \n", - "9 0 /content/datasets/brain-tumor/images/train/000... \n", - "10 0 /content/datasets/brain-tumor/images/train/000... \n", - "11 0 /content/datasets/brain-tumor/images/train/000... \n", - "12 0 /content/datasets/brain-tumor/images/train/000... \n", - "13 0 /content/datasets/brain-tumor/images/train/000... \n", - "14 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - "15 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - "16 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - "17 0 /content/datasets/brain-tumor/images/train/000... \n", - "18 0 /content/datasets/brain-tumor/images/train/000... \n", - "19 0 /content/datasets/brain-tumor/images/train/000... \n", - "20 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - " 3 None \n", - "21 0 /content/datasets/brain-tumor/images/train/000... \n", - "22 0 /content/datasets/brain-tumor/images/train/000... \n", - "23 0 /content/datasets/brain-tumor/images/train/000... \n", - "24 0 /content/datasets/brain-tumor/images/train/000... \n", - "25 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - "26 0 /content/datasets/brain-tumor/images/train/000... \n", - "27 0 /content/datasets/brain-tumor/images/train/000... \n", - "28 0 /content/datasets/brain-tumor/images/train/000... \n", - "29 0 /content/datasets/brain-tumor/images/train/000... \n", - "30 0 /content/datasets/brain-tumor/images/train/000... \n", - "31 0 /content/datasets/brain-tumor/images/train/000... \n", - "32 0 /content/datasets/brain-tumor/images/train/000... \n", - "33 0 /content/datasets/brain-tumor/images/train/000... \n", - "34 0 /content/datasets/brain-tumor/images/train/000... \n", - "35 0 /content/datasets/brain-tumor/images/train/000... \n", - "36 0 /content/datasets/brain-tumor/images/train/000... \n", - "\n", - " cls \n", - "sample_id annotation_id \n", - "0 0 [[1.0]] \n", - "1 0 [[1.0]] \n", - "2 0 [[1.0], [1.0]] \n", - " 1 None \n", - " 2 None \n", - "3 0 [[1.0]] \n", - "4 0 [[1.0]] \n", - "5 0 [[1.0]] \n", - "6 0 [[1.0]] \n", - "7 0 [[1.0]] \n", - "8 0 [[1.0]] \n", - "9 0 [[1.0]] \n", - "10 0 [[1.0]] \n", - "11 0 [[1.0]] \n", - "12 0 [[1.0]] \n", - "13 0 [[1.0]] \n", - "14 0 [[1.0], [1.0]] \n", - " 1 None \n", - " 2 None \n", - "15 0 [[1.0], [1.0]] \n", - " 1 None \n", - " 2 None \n", - "16 0 [[1.0], [1.0]] \n", - " 1 None \n", - " 2 None \n", - "17 0 [[1.0]] \n", - "18 0 [[1.0]] \n", - "19 0 [[0.0]] \n", - "20 0 [[0.0], [0.0], [0.0]] \n", - " 1 None \n", - " 2 None \n", - " 3 None \n", - "21 0 [[0.0]] \n", - "22 0 [[0.0]] \n", - "23 0 [[1.0]] \n", - "24 0 [[1.0]] \n", - "25 0 [[0.0], [0.0]] \n", - " 1 None \n", - " 2 None \n", - "26 0 [[0.0]] \n", - "27 0 [[0.0]] \n", - "28 0 [[1.0]] \n", - "29 0 [[1.0]] \n", - "30 0 [[1.0]] \n", - "31 0 [[1.0]] \n", - "32 0 [[1.0]] \n", - "33 0 [[1.0]] \n", - "34 0 [[1.0]] \n", - "35 0 [[1.0]] \n", - "36 0 [[1.0]] " - ] - }, - "metadata": {}, - "output_type": "display_data" - }, - { - "data": { - "text/html": [ - "🔗 Open the full grid in Weights Studio" - ], - "text/plain": [ - "" - ] - }, - "metadata": {}, - "output_type": "display_data" - } - ], + "outputs": [], "source": [ "import pandas as pd\n", "from IPython.display import display, HTML\n", diff --git a/weightslab/examples/Notebooks/Ultralytics/wl-how-to-train-ultralytics-yolo-on-kitti-detection-dataset.ipynb b/weightslab/examples/Notebooks/Ultralytics/wl-how-to-train-ultralytics-yolo-on-kitti-detection-dataset.ipynb index 0531b804..a3c408d6 100644 --- a/weightslab/examples/Notebooks/Ultralytics/wl-how-to-train-ultralytics-yolo-on-kitti-detection-dataset.ipynb +++ b/weightslab/examples/Notebooks/Ultralytics/wl-how-to-train-ultralytics-yolo-on-kitti-detection-dataset.ipynb @@ -44,7 +44,7 @@ }, { "cell_type": "code", - "execution_count": 6, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/" @@ -52,24 +52,14 @@ "id": "71be394c", "outputId": "c9ef4767-84f5-4bea-e395-9e11766002f2" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Requirement already satisfied: setuptools in /usr/local/lib/python3.12/dist-packages (75.2.0)\n", - "Requirement already satisfied: wheel in /usr/local/lib/python3.12/dist-packages (0.47.0)\n", - "Requirement already satisfied: packaging>=24.0 in /usr/local/lib/python3.12/dist-packages (from wheel) (26.2)\n" - ] - } - ], + "outputs": [], "source": [ "%pip install setuptools wheel" ] }, { "cell_type": "code", - "execution_count": 7, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/", @@ -78,149 +68,7 @@ "id": "jV1oi_PQN0xK", "outputId": "eda8eca1-f4b3-4ffd-eaa6-0f67ce857153" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Collecting git+https://github.com/GrayboxTech/weightslab.git\n", - " Cloning https://github.com/GrayboxTech/weightslab.git (to revision landingcollab) to /tmp/pip-req-build-z4ba7cx9\n", - " Running command git clone --filter=blob:none --quiet https://github.com/GrayboxTech/weightslab.git /tmp/pip-req-build-z4ba7cx9\n", - " Running command git checkout -b landingcollab --track origin/landingcollab\n", - " Switched to a new branch 'landingcollab'\n", - " Branch 'landingcollab' set up to track remote branch 'landingcollab' from 'origin'.\n", - " Resolved https://github.com/GrayboxTech/weightslab.git to commit 3f1cf143bc93e67df43adc71ab831fa71477e0d2\n", - " Installing build dependencies ... \u001b[?25l\u001b[?25hdone\n", - " Getting requirements to build wheel ... \u001b[?25l\u001b[?25hdone\n", - " Preparing metadata (pyproject.toml) ... \u001b[?25l\u001b[?25hdone\n", - "Requirement already satisfied: numpy<3,>=1.24 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (2.0.2)\n", - "Requirement already satisfied: pandas<3,>=2.2.2 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (2.2.2)\n", - "Requirement already satisfied: duckdb<2,>=1.1 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.3.2)\n", - "Requirement already satisfied: PyYAML<7,>=6.0.3 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (6.0.3)\n", - "Requirement already satisfied: dill<0.5,>=0.3.8 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (0.3.8)\n", - "Requirement already satisfied: zstandard<1,>=0.22 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (0.25.0)\n", - "Requirement already satisfied: h5py<4,>=3.10 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (3.16.0)\n", - "Requirement already satisfied: xxhash<4.1,>=3.4 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (3.7.1)\n", - "Requirement already satisfied: tables<4,>=3.9 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (3.10.2)\n", - "Requirement already satisfied: torch<=2.9,>=2.1 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (2.9.0)\n", - "Requirement already satisfied: torchvision<1,>=0.16 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (0.24.0)\n", - "Requirement already satisfied: torchmetrics>=1.9 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.9.0)\n", - "Requirement already satisfied: grpcio<2,>=1.80 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.81.1)\n", - "Requirement already satisfied: protobuf<8,>=5.28.1 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (5.29.6)\n", - "Requirement already satisfied: pydantic<3,>=2.7 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (2.13.4)\n", - "Requirement already satisfied: Pillow<12,>=10 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (11.3.0)\n", - "Requirement already satisfied: graphviz<1,>=0.20 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (0.21)\n", - "Requirement already satisfied: onnx<=1.20,>=1.15 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.20.0)\n", - "Requirement already satisfied: tqdm<5,>=4.66 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (4.67.3)\n", - "Requirement already satisfied: python-dotenv<2,>=1 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.2.2)\n", - "Requirement already satisfied: langchain-core<2,>=0.3 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.4.9)\n", - "Requirement already satisfied: langchain-ollama<2,>=0.2 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.1.0)\n", - "Requirement already satisfied: langchain-openai<2,>=0.2 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.3.5)\n", - "Requirement already satisfied: typing-extensions~=4.12 in /usr/local/lib/python3.12/dist-packages (from grpcio<2,>=1.80->weightslab==1.3.3.dev8) (4.15.0)\n", - "Requirement already satisfied: jsonpatch<2.0.0,>=1.33.0 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (1.33)\n", - "Requirement already satisfied: langchain-protocol>=0.0.17 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (0.0.18)\n", - "Requirement already satisfied: langsmith<1.0.0,>=0.3.45 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (0.9.1)\n", - "Requirement already satisfied: packaging>=23.2.0 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (26.2)\n", - "Requirement already satisfied: tenacity!=8.4.0,<10.0.0,>=8.1.0 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (9.1.4)\n", - "Requirement already satisfied: uuid-utils<1.0,>=0.12.0 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (0.16.2)\n", - "Requirement already satisfied: ollama<1.0.0,>=0.6.1 in /usr/local/lib/python3.12/dist-packages (from langchain-ollama<2,>=0.2->weightslab==1.3.3.dev8) (0.6.2)\n", - "Requirement already satisfied: openai<3.0.0,>=2.45.0 in /usr/local/lib/python3.12/dist-packages (from langchain-openai<2,>=0.2->weightslab==1.3.3.dev8) (2.45.0)\n", - "Requirement already satisfied: tiktoken<1.0.0,>=0.7.0 in /usr/local/lib/python3.12/dist-packages (from langchain-openai<2,>=0.2->weightslab==1.3.3.dev8) (0.13.0)\n", - "Requirement already satisfied: ml_dtypes>=0.5.0 in /usr/local/lib/python3.12/dist-packages (from onnx<=1.20,>=1.15->weightslab==1.3.3.dev8) (0.5.4)\n", - "Requirement already satisfied: python-dateutil>=2.8.2 in /usr/local/lib/python3.12/dist-packages (from pandas<3,>=2.2.2->weightslab==1.3.3.dev8) (2.9.0.post0)\n", - "Requirement already satisfied: pytz>=2020.1 in /usr/local/lib/python3.12/dist-packages (from pandas<3,>=2.2.2->weightslab==1.3.3.dev8) (2025.2)\n", - "Requirement already satisfied: tzdata>=2022.7 in /usr/local/lib/python3.12/dist-packages (from pandas<3,>=2.2.2->weightslab==1.3.3.dev8) (2026.2)\n", - "Requirement already satisfied: annotated-types>=0.6.0 in /usr/local/lib/python3.12/dist-packages (from pydantic<3,>=2.7->weightslab==1.3.3.dev8) (0.7.0)\n", - "Requirement already satisfied: pydantic-core==2.46.4 in /usr/local/lib/python3.12/dist-packages (from pydantic<3,>=2.7->weightslab==1.3.3.dev8) (2.46.4)\n", - "Requirement already satisfied: typing-inspection>=0.4.2 in /usr/local/lib/python3.12/dist-packages (from pydantic<3,>=2.7->weightslab==1.3.3.dev8) (0.4.2)\n", - "Requirement already satisfied: numexpr>=2.6.2 in /usr/local/lib/python3.12/dist-packages (from tables<4,>=3.9->weightslab==1.3.3.dev8) (2.14.1)\n", - "Requirement already satisfied: py-cpuinfo in /usr/local/lib/python3.12/dist-packages (from tables<4,>=3.9->weightslab==1.3.3.dev8) (9.0.0)\n", - "Requirement already satisfied: blosc2>=2.3.0 in /usr/local/lib/python3.12/dist-packages (from tables<4,>=3.9->weightslab==1.3.3.dev8) (4.5.1)\n", - "Requirement already satisfied: filelock in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.29.4)\n", - "Requirement already satisfied: setuptools in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (75.2.0)\n", - "Requirement already satisfied: sympy>=1.13.3 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (1.14.0)\n", - "Requirement already satisfied: networkx>=2.5.1 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.6.1)\n", - "Requirement already satisfied: jinja2 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.1.6)\n", - "Requirement already satisfied: fsspec>=0.8.5 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (2025.3.0)\n", - "Requirement already satisfied: nvidia-cuda-nvrtc-cu12==12.8.93 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.93)\n", - "Requirement already satisfied: nvidia-cuda-runtime-cu12==12.8.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.90)\n", - "Requirement already satisfied: nvidia-cuda-cupti-cu12==12.8.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.90)\n", - "Requirement already satisfied: nvidia-cudnn-cu12==9.10.2.21 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (9.10.2.21)\n", - "Requirement already satisfied: nvidia-cublas-cu12==12.8.4.1 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.4.1)\n", - "Requirement already satisfied: nvidia-cufft-cu12==11.3.3.83 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (11.3.3.83)\n", - "Requirement already satisfied: nvidia-curand-cu12==10.3.9.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (10.3.9.90)\n", - "Requirement already satisfied: nvidia-cusolver-cu12==11.7.3.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (11.7.3.90)\n", - "Requirement already satisfied: nvidia-cusparse-cu12==12.5.8.93 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.5.8.93)\n", - "Requirement already satisfied: nvidia-cusparselt-cu12==0.7.1 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (0.7.1)\n", - "Requirement already satisfied: nvidia-nccl-cu12==2.27.5 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (2.27.5)\n", - "Requirement already satisfied: nvidia-nvshmem-cu12==3.3.20 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.3.20)\n", - "Requirement already satisfied: nvidia-nvtx-cu12==12.8.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.90)\n", - "Requirement already satisfied: nvidia-nvjitlink-cu12==12.8.93 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.93)\n", - "Requirement already satisfied: nvidia-cufile-cu12==1.13.1.3 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (1.13.1.3)\n", - "Requirement already satisfied: triton==3.5.0 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.5.0)\n", - "Requirement already satisfied: lightning-utilities>=0.15.3 in /usr/local/lib/python3.12/dist-packages (from torchmetrics>=1.9->weightslab==1.3.3.dev8) (0.15.3)\n", - "Requirement already satisfied: ndindex in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (1.10.1)\n", - "Requirement already satisfied: msgpack in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (1.2.1)\n", - "Requirement already satisfied: requests in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (2.32.4)\n", - "Requirement already satisfied: rich in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (13.9.4)\n", - "Requirement already satisfied: textual in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (6.2.1)\n", - "Requirement already satisfied: textual-plotext in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (1.0.1)\n", - "Requirement already satisfied: threadpoolctl in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (3.6.0)\n", - "Requirement already satisfied: jsonpointer>=1.9 in /usr/local/lib/python3.12/dist-packages (from jsonpatch<2.0.0,>=1.33.0->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (3.1.1)\n", - "Requirement already satisfied: anyio>=3.5.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (4.14.0)\n", - "Requirement already satisfied: distro>=1.7.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (1.9.0)\n", - "Requirement already satisfied: httpx<1,>=0.23.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (0.28.1)\n", - "Requirement already satisfied: orjson>=3.9.14 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (3.11.9)\n", - "Requirement already satisfied: requests-toolbelt>=1.0.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (1.0.0)\n", - "Requirement already satisfied: sniffio>=1.1 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (1.3.1)\n", - "Requirement already satisfied: websockets>=15.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (15.0.1)\n", - "Requirement already satisfied: jiter<1,>=0.10.0 in /usr/local/lib/python3.12/dist-packages (from openai<3.0.0,>=2.45.0->langchain-openai<2,>=0.2->weightslab==1.3.3.dev8) (0.15.0)\n", - "Requirement already satisfied: six>=1.5 in /usr/local/lib/python3.12/dist-packages (from python-dateutil>=2.8.2->pandas<3,>=2.2.2->weightslab==1.3.3.dev8) (1.17.0)\n", - "Requirement already satisfied: mpmath<1.4,>=1.1.0 in /usr/local/lib/python3.12/dist-packages (from sympy>=1.13.3->torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (1.3.0)\n", - "Requirement already satisfied: regex in /usr/local/lib/python3.12/dist-packages (from tiktoken<1.0.0,>=0.7.0->langchain-openai<2,>=0.2->weightslab==1.3.3.dev8) (2025.11.3)\n", - "Requirement already satisfied: MarkupSafe>=2.0 in /usr/local/lib/python3.12/dist-packages (from jinja2->torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.0.3)\n", - "Requirement already satisfied: idna>=2.8 in /usr/local/lib/python3.12/dist-packages (from anyio>=3.5.0->langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (3.18)\n", - "Requirement already satisfied: certifi in /usr/local/lib/python3.12/dist-packages (from httpx<1,>=0.23.0->langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (2026.6.17)\n", - "Requirement already satisfied: httpcore==1.* in /usr/local/lib/python3.12/dist-packages (from httpx<1,>=0.23.0->langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (1.0.9)\n", - "Requirement already satisfied: h11>=0.16 in /usr/local/lib/python3.12/dist-packages (from httpcore==1.*->httpx<1,>=0.23.0->langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (0.16.0)\n", - "Requirement already satisfied: charset_normalizer<4,>=2 in /usr/local/lib/python3.12/dist-packages (from requests->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (3.4.7)\n", - "Requirement already satisfied: urllib3<3,>=1.21.1 in /usr/local/lib/python3.12/dist-packages (from requests->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (2.5.0)\n", - "Requirement already satisfied: markdown-it-py>=2.2.0 in /usr/local/lib/python3.12/dist-packages (from rich->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (4.2.0)\n", - "Requirement already satisfied: pygments<3.0.0,>=2.13.0 in /usr/local/lib/python3.12/dist-packages (from rich->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (2.20.0)\n", - "Requirement already satisfied: platformdirs<5,>=3.6.0 in /usr/local/lib/python3.12/dist-packages (from textual->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (4.10.0)\n", - "Requirement already satisfied: plotext<6.0.0,>=5.2.8 in /usr/local/lib/python3.12/dist-packages (from textual-plotext->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (5.3.2)\n", - "Requirement already satisfied: mdurl~=0.1 in /usr/local/lib/python3.12/dist-packages (from markdown-it-py>=2.2.0->rich->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (0.1.2)\n", - "Requirement already satisfied: linkify-it-py<3,>=1 in /usr/local/lib/python3.12/dist-packages (from markdown-it-py[linkify,plugins]>=2.1.0->textual->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (2.1.0)\n", - "Requirement already satisfied: mdit-py-plugins>=0.5.0 in /usr/local/lib/python3.12/dist-packages (from markdown-it-py[linkify,plugins]>=2.1.0->textual->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (0.6.1)\n", - "Requirement already satisfied: uc-micro-py in /usr/local/lib/python3.12/dist-packages (from linkify-it-py<3,>=1->markdown-it-py[linkify,plugins]>=2.1.0->textual->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (2.0.0)\n", - "Building wheels for collected packages: weightslab\n", - " Building wheel for weightslab (pyproject.toml) ... \u001b[?25l\u001b[?25hdone\n", - " Created wheel for weightslab: filename=weightslab-1.3.3.dev8-py3-none-any.whl size=2827715 sha256=5c4e7c14bbf95a48b9a8c86b5a8c744af67e1eb395197cbe79d3e443c21cf86e\n", - " Stored in directory: /tmp/pip-ephem-wheel-cache-__zwgf7x/wheels/4d/c7/cd/817c2745790555e7b6c532f5b6055d47567f9416fdfc65879b\n", - "Successfully built weightslab\n", - "Installing collected packages: weightslab\n", - " Attempting uninstall: weightslab\n", - " Found existing installation: weightslab 1.3.3.dev7\n", - " Uninstalling weightslab-1.3.3.dev7:\n", - " Successfully uninstalled weightslab-1.3.3.dev7\n", - "Successfully installed weightslab-1.3.3.dev8\n" - ] - }, - { - "data": { - "application/vnd.colab-display-data+json": { - "id": "b7fce274042740d19087972fe9c01c9b", - "pip_warning": { - "packages": [ - "weightslab" - ] - } - } - }, - "metadata": {}, - "output_type": "display_data" - } - ], + "outputs": [], "source": [ "%pip install weightslab\n", "!pip install --upgrade \"protobuf>=6.31.1\" # Google Collab compat update" @@ -228,7 +76,7 @@ }, { "cell_type": "code", - "execution_count": 1, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/" @@ -236,39 +84,7 @@ "id": "pJhvREbaKs7f", "outputId": "651baa36-9ebb-49da-830b-e6c8a6e2b64f" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Ultralytics 8.4.96 🚀 Python-3.12.13 torch-2.9.0+cu128 CPU (Intel Xeon CPU @ 2.20GHz)\n", - "Setup complete ✅ (2 CPUs, 12.7 GB RAM, 30.5/107.7 GB disk)\n", - "16/07/2026-09:29:29.683 INFO:root:setup_logging: WeightsLab logging initialized - Log file: /tmp/tmp6meuzqdz/weightslab_logs/weightslab_20260716_092929.log\n", - "16/07/2026-09:29:29.685 INFO:weightslab:: WeightsLab package initialized - Log level: INFO, Log to file: True\n", - "16/07/2026-09:29:29.686 INFO:weightslab:: \n", - "╭─────────────────────────────────────────────────────────────────────────────────────────────────────╮\n", - "│ │\n", - "│ \u001b[32m$$\\ $$\\ \u001b[0m $$\\ $$\\ $$\\ \u001b[31m$$\\ \u001b[0m $$\\ │\n", - "│ \u001b[32m$$ | $\\ $$ |\u001b[0m \\__| $$ | $$ | \u001b[31m$$ | \u001b[0m $$ | │\n", - "│ \u001b[32m$$ |$$$\\ $$ |\u001b[0m $$$$$$\\ $$\\ $$$$$$\\ $$$$$$$\\ $$$$$$\\ $$$$$$$\\ \u001b[31m$$ | \u001b[0m $$$$$$\\ $$$$$$$\\ │\n", - "│ \u001b[32m$$ $$ $$\\$$ |\u001b[0m$$ __$$\\ $$ |$$ __$$\\ $$ __$$\\ \\_$$ _| $$ _____|\u001b[31m$$ | \u001b[0m \\____$$\\ $$ __$$\\ │\n", - "│ \u001b[32m$$$$ _$$$$ |\u001b[0m$$$$$$$$ |$$ |$$ / $$ |$$ | $$ | $$ | \\$$$$$$\\ \u001b[31m$$ | \u001b[0m $$$$$$$ |$$ | $$ | │\n", - "│ \u001b[32m$$$ / \\$$$ |\u001b[0m$$ ____|$$ |$$ | $$ |$$ | $$ | $$ |$$\\ \\____$$\\ \u001b[31m$$ | \u001b[0m$$ __$$ |$$ | $$ | │\n", - "│ \u001b[32m$$ / \\$$ |\u001b[0m\\$$$$$$$\\ $$ |\\$$$$$$$ |$$ | $$ | \\$$$$ |$$$$$$$ |\u001b[31m$$$$$$$$\\ \u001b[0m\\$$$$$$$ |$$$$$$$ | │\n", - "│ \u001b[32m\\__/ \\__|\u001b[0m \\_______|\\__| \\____$$ |\\__| \\__| \\____/ \\_______/ \u001b[31m\\________|\u001b[0m \\_______|\\_______/ │\n", - "│ \u001b[32m \u001b[0m $$\\ $$ |\u001b[31m\u001b[0m │\n", - "│ \u001b[32m \u001b[0m \\$$$$$$ |\u001b[31m\u001b[0m │\n", - "│ \u001b[32m \u001b[0m \\______/\u001b[31m\u001b[0m │\n", - "│ │\n", - "│ Inspect - Edit - Evolve Neural Networks │\n", - "│ │\n", - "╰─── By GrayBx, v1.3.3.dev8 ──────────────────────────────────────────────────────────────────────────╯\n", - "\n", - "\n", - "16/07/2026-09:29:29.690 INFO:weightslab.utils.telemetry:ping_import: WeightsLab uses anonymous usage data (package version used, and OS name). Set WL_NO_TELEMETRY=1 to disable.\n" - ] - } - ], + "outputs": [], "source": [ "# Install WeightsLab. Colab already ships torch, torchvision, numpy, scikit-learn\n", "# and Pillow, so nothing extra is needed here.\n", @@ -356,7 +172,6 @@ "import warnings\n", "warnings.filterwarnings(\"ignore\")\n", "\n", - "from weightslab.integrations.ultralytics import WLAwareTrainer\n", "from ultralytics import YOLO\n", "\n", "# Hyperparameters registered with WeightsLab -> live-editable from Weights Studio.\n", @@ -381,7 +196,7 @@ "# Read back the (now live) hyperparameters and hand training to WLAwareTrainer.\n", "model = YOLO(cfg[\"model\"])\n", "results = model.train(\n", - " trainer=WLAwareTrainer,\n", + " trainer=wl.WLAwareTrainer,\n", " data=str(cfg[\"data\"]),\n", " imgsz=cfg[\"image_size\"],\n", " epochs=int(cfg[\"epochs\"]),\n", @@ -411,7 +226,7 @@ }, { "cell_type": "code", - "execution_count": 3, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/", @@ -421,1230 +236,7 @@ "id": "uOJdXGA3Ks7o", "outputId": "26652fa3-6a15-43ce-9d75-034eaba21ccd" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "1241 samples tracked\n" - ] - }, - { - "data": { - "application/vnd.google.colaboratory.intrinsic+json": { - "repr_error": "Out of range float values are not JSON compliant: nan", - "type": "dataframe" - }, - "text/html": [ - "\n", - "
\n", - "
\n", - "\n", - "\n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - "
predictionprediction_rawtargetdiscardedorigintask_typelast_seengroup_idmember_rankimg_pathcls
sample_idannotation_id
00NoneNone[0.2335685, 0.254695, 0.4553995, 0.43075103, 1...Falsetrain_loader-1.000.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
10NoneNone[0.251174, 0.24061048, 0.44366202, 0.4307515, ...Falsetrain_loader-1.010.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
20NoneNoneNoneFalsetrain_loader-1.020.0/content/datasets/brain-tumor/images/train/000...[[1.0], [1.0]]
1NoneNone[0.523474, 0.3978875, 0.634976, 0.48122054, 1....NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.40610352, 0.415493, 0.5234745, 0.522301, 1....NoneNaNNoneNaNNoneNaNNoneNone
30NoneNone[0.403756, 0.3826295, 0.637324, 0.5152585, 1.0...Falsetrain_loader-1.030.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
40NoneNone[0.4436615, 0.3814555, 0.5938965, 0.4518785, 1...Falsetrain_loader-1.040.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
50NoneNone[0.6044605, 0.3990615, 0.67253554, 0.46596247,...Falsetrain_loader-1.050.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
60NoneNone[0.410798, 0.4436625, 0.556338, 0.51173747, 1....Falsetrain_loader-1.060.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
70NoneNone[0.650235, 0.4166665, 0.73826295, 0.48356748, ...Falsetrain_loader-1.070.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
80NoneNone[0.64788747, 0.388498, 0.7582165, 0.49413198, ...Falsetrain_loader-1.080.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
90NoneNone[0.6807515, 0.40140802, 0.75704247, 0.45774597...Falsetrain_loader-1.090.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
100NoneNone[0.431925, 0.1701875, 0.561033, 0.3180745, 1.0...Falsetrain_loader-1.0100.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
110NoneNone[0.42018852, 0.34507048, 0.5434275, 0.4941315,...Falsetrain_loader-1.0110.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
120NoneNone[0.3732395, 0.348592, 0.54812247, 0.46009403, ...Falsetrain_loader-1.0120.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
130NoneNone[0.37793398, 0.3814555, 0.48122, 0.45892054, 1...Falsetrain_loader-1.0130.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
140NoneNoneNoneFalsetrain_loader-1.0140.0/content/datasets/brain-tumor/images/train/000...[[1.0], [1.0]]
1NoneNone[0.43427247, 0.3556335, 0.51173747, 0.43075046...NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.4929575, 0.44718304, 0.5375585, 0.485915, 1...NoneNaNNoneNaNNoneNaNNoneNone
150NoneNoneNoneFalsetrain_loader-1.0150.0/content/datasets/brain-tumor/images/train/000...[[1.0], [1.0]]
1NoneNone[0.629108, 0.39788702, 0.67840403, 0.453051, 1...NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.4530515, 0.4107985, 0.58568054, 0.5105635, ...NoneNaNNoneNaNNoneNaNNoneNone
160NoneNoneNoneFalsetrain_loader-1.0160.0/content/datasets/brain-tumor/images/train/000...[[1.0], [1.0]]
1NoneNone[0.45070451, 0.3967135, 0.5762915, 0.5152585, ...NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.63732404, 0.403756, 0.684272, 0.46009403, 1...NoneNaNNoneNaNNoneNaNNoneNone
170NoneNone[0.43779355, 0.3826295, 0.61032856, 0.5223005,...Falsetrain_loader-1.0170.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
180NoneNone[0.4401405, 0.374413, 0.6068075, 0.529343, 1.0...Falsetrain_loader-1.0180.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
190NoneNone[0.536385, 0.4495305, 0.600939, 0.49999946, 0....Falsetrain_loader-1.0190.0/content/datasets/brain-tumor/images/train/000...[[0.0]]
200NoneNoneNoneFalsetrain_loader-1.0200.0/content/datasets/brain-tumor/images/train/000...[[0.0], [0.0], [0.0]]
1NoneNone[0.52347445, 0.4424885, 0.6150235, 0.5281695, ...NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.456573, 0.515258, 0.504695, 0.562206, 0.0, ...NoneNaNNoneNaNNoneNaNNoneNone
3NoneNone[0.504695, 0.44248852, 0.54342705, 0.48004752,...NoneNaNNoneNaNNoneNaNNoneNone
210NoneNone[0.54107946, 0.456573, 0.6103285, 0.511737, 0....Falsetrain_loader-1.0210.0/content/datasets/brain-tumor/images/train/000...[[0.0]]
220NoneNone[0.3298125, 0.2042255, 0.4225355, 0.2957745, 0...Falsetrain_loader-1.0220.0/content/datasets/brain-tumor/images/train/000...[[0.0]]
230NoneNone[0.517606, 0.18661949, 0.615024, 0.2875585, 1....Falsetrain_loader-1.0230.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
240NoneNone[0.51173747, 0.18779299, 0.6126765, 0.308685, ...Falsetrain_loader-1.0240.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
250NoneNoneNoneFalsetrain_loader-1.0250.0/content/datasets/brain-tumor/images/train/000...[[0.0], [0.0]]
1NoneNone[0.30985948, 0.3626765, 0.3861505, 0.43896753,...NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.41314548, 0.3427225, 0.48122054, 0.42018753...NoneNaNNoneNaNNoneNaNNoneNone
260NoneNone[0.25821602, 0.30046952, 0.46830997, 0.4577465...Falsetrain_loader-1.0260.0/content/datasets/brain-tumor/images/train/000...[[0.0]]
270NoneNone[0.2734745, 0.2828635, 0.44718352, 0.4577465, ...Falsetrain_loader-1.0270.0/content/datasets/brain-tumor/images/train/000...[[0.0]]
280NoneNone[0.5915495, 0.23591551, 0.6455405, 0.2887325, ...Falsetrain_loader-1.0280.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
290NoneNone[0.5903755, 0.231221, 0.6420185, 0.280517, 1.0...Falsetrain_loader-1.0290.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
300NoneNone[0.40493003, 0.2018775, 0.469484, 0.2699525, 1...Falsetrain_loader-1.0300.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
310NoneNone[0.39554, 0.19483551, 0.483568, 0.2887325, 1.0...Falsetrain_loader-1.0310.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
320NoneNone[0.4025825, 0.212441, 0.45539945, 0.269953, 1....Falsetrain_loader-1.0320.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
330NoneNone[0.2265255, 0.2136145, 0.3438965, 0.30751148, ...Falsetrain_loader-1.0330.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
340NoneNone[0.634976, 0.349765, 0.7723, 0.456573, 1.0, 1.0]Falsetrain_loader-1.0340.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
350NoneNone[0.6079815, 0.320423, 0.7969485, 0.489437, 1.0...Falsetrain_loader-1.0350.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
360NoneNone[0.6185445, 0.30164298, 0.7758215, 0.494131, 1...Falsetrain_loader-1.0360.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
\n", - "
\n", - "
\n", - "\n", - "
\n", - " \n", - "\n", - " \n", - "\n", - " \n", - "
\n", - "\n", - "\n", - "
\n", - "
\n" - ], - "text/plain": [ - " prediction prediction_raw \\\n", - "sample_id annotation_id \n", - "0 0 None None \n", - "1 0 None None \n", - "2 0 None None \n", - " 1 None None \n", - " 2 None None \n", - "3 0 None None \n", - "4 0 None None \n", - "5 0 None None \n", - "6 0 None None \n", - "7 0 None None \n", - "8 0 None None \n", - "9 0 None None \n", - "10 0 None None \n", - "11 0 None None \n", - "12 0 None None \n", - "13 0 None None \n", - "14 0 None None \n", - " 1 None None \n", - " 2 None None \n", - "15 0 None None \n", - " 1 None None \n", - " 2 None None \n", - "16 0 None None \n", - " 1 None None \n", - " 2 None None \n", - "17 0 None None \n", - "18 0 None None \n", - "19 0 None None \n", - "20 0 None None \n", - " 1 None None \n", - " 2 None None \n", - " 3 None None \n", - "21 0 None None \n", - "22 0 None None \n", - "23 0 None None \n", - "24 0 None None \n", - "25 0 None None \n", - " 1 None None \n", - " 2 None None \n", - "26 0 None None \n", - "27 0 None None \n", - "28 0 None None \n", - "29 0 None None \n", - "30 0 None None \n", - "31 0 None None \n", - "32 0 None None \n", - "33 0 None None \n", - "34 0 None None \n", - "35 0 None None \n", - "36 0 None None \n", - "\n", - " target \\\n", - "sample_id annotation_id \n", - "0 0 [0.2335685, 0.254695, 0.4553995, 0.43075103, 1... \n", - "1 0 [0.251174, 0.24061048, 0.44366202, 0.4307515, ... \n", - "2 0 None \n", - " 1 [0.523474, 0.3978875, 0.634976, 0.48122054, 1.... \n", - " 2 [0.40610352, 0.415493, 0.5234745, 0.522301, 1.... \n", - "3 0 [0.403756, 0.3826295, 0.637324, 0.5152585, 1.0... \n", - "4 0 [0.4436615, 0.3814555, 0.5938965, 0.4518785, 1... \n", - "5 0 [0.6044605, 0.3990615, 0.67253554, 0.46596247,... \n", - "6 0 [0.410798, 0.4436625, 0.556338, 0.51173747, 1.... \n", - "7 0 [0.650235, 0.4166665, 0.73826295, 0.48356748, ... \n", - "8 0 [0.64788747, 0.388498, 0.7582165, 0.49413198, ... \n", - "9 0 [0.6807515, 0.40140802, 0.75704247, 0.45774597... \n", - "10 0 [0.431925, 0.1701875, 0.561033, 0.3180745, 1.0... \n", - "11 0 [0.42018852, 0.34507048, 0.5434275, 0.4941315,... \n", - "12 0 [0.3732395, 0.348592, 0.54812247, 0.46009403, ... \n", - "13 0 [0.37793398, 0.3814555, 0.48122, 0.45892054, 1... \n", - "14 0 None \n", - " 1 [0.43427247, 0.3556335, 0.51173747, 0.43075046... \n", - " 2 [0.4929575, 0.44718304, 0.5375585, 0.485915, 1... \n", - "15 0 None \n", - " 1 [0.629108, 0.39788702, 0.67840403, 0.453051, 1... \n", - " 2 [0.4530515, 0.4107985, 0.58568054, 0.5105635, ... \n", - "16 0 None \n", - " 1 [0.45070451, 0.3967135, 0.5762915, 0.5152585, ... \n", - " 2 [0.63732404, 0.403756, 0.684272, 0.46009403, 1... \n", - "17 0 [0.43779355, 0.3826295, 0.61032856, 0.5223005,... \n", - "18 0 [0.4401405, 0.374413, 0.6068075, 0.529343, 1.0... \n", - "19 0 [0.536385, 0.4495305, 0.600939, 0.49999946, 0.... \n", - "20 0 None \n", - " 1 [0.52347445, 0.4424885, 0.6150235, 0.5281695, ... \n", - " 2 [0.456573, 0.515258, 0.504695, 0.562206, 0.0, ... \n", - " 3 [0.504695, 0.44248852, 0.54342705, 0.48004752,... \n", - "21 0 [0.54107946, 0.456573, 0.6103285, 0.511737, 0.... \n", - "22 0 [0.3298125, 0.2042255, 0.4225355, 0.2957745, 0... \n", - "23 0 [0.517606, 0.18661949, 0.615024, 0.2875585, 1.... \n", - "24 0 [0.51173747, 0.18779299, 0.6126765, 0.308685, ... \n", - "25 0 None \n", - " 1 [0.30985948, 0.3626765, 0.3861505, 0.43896753,... \n", - " 2 [0.41314548, 0.3427225, 0.48122054, 0.42018753... \n", - "26 0 [0.25821602, 0.30046952, 0.46830997, 0.4577465... \n", - "27 0 [0.2734745, 0.2828635, 0.44718352, 0.4577465, ... \n", - "28 0 [0.5915495, 0.23591551, 0.6455405, 0.2887325, ... \n", - "29 0 [0.5903755, 0.231221, 0.6420185, 0.280517, 1.0... \n", - "30 0 [0.40493003, 0.2018775, 0.469484, 0.2699525, 1... \n", - "31 0 [0.39554, 0.19483551, 0.483568, 0.2887325, 1.0... \n", - "32 0 [0.4025825, 0.212441, 0.45539945, 0.269953, 1.... \n", - "33 0 [0.2265255, 0.2136145, 0.3438965, 0.30751148, ... \n", - "34 0 [0.634976, 0.349765, 0.7723, 0.456573, 1.0, 1.0] \n", - "35 0 [0.6079815, 0.320423, 0.7969485, 0.489437, 1.0... \n", - "36 0 [0.6185445, 0.30164298, 0.7758215, 0.494131, 1... \n", - "\n", - " discarded origin task_type last_seen group_id \\\n", - "sample_id annotation_id \n", - "0 0 False train_loader -1.0 0 \n", - "1 0 False train_loader -1.0 1 \n", - "2 0 False train_loader -1.0 2 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - "3 0 False train_loader -1.0 3 \n", - "4 0 False train_loader -1.0 4 \n", - "5 0 False train_loader -1.0 5 \n", - "6 0 False train_loader -1.0 6 \n", - "7 0 False train_loader -1.0 7 \n", - "8 0 False train_loader -1.0 8 \n", - "9 0 False train_loader -1.0 9 \n", - "10 0 False train_loader -1.0 10 \n", - "11 0 False train_loader -1.0 11 \n", - "12 0 False train_loader -1.0 12 \n", - "13 0 False train_loader -1.0 13 \n", - "14 0 False train_loader -1.0 14 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - "15 0 False train_loader -1.0 15 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - "16 0 False train_loader -1.0 16 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - "17 0 False train_loader -1.0 17 \n", - "18 0 False train_loader -1.0 18 \n", - "19 0 False train_loader -1.0 19 \n", - "20 0 False train_loader -1.0 20 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - " 3 None NaN None NaN None \n", - "21 0 False train_loader -1.0 21 \n", - "22 0 False train_loader -1.0 22 \n", - "23 0 False train_loader -1.0 23 \n", - "24 0 False train_loader -1.0 24 \n", - "25 0 False train_loader -1.0 25 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - "26 0 False train_loader -1.0 26 \n", - "27 0 False train_loader -1.0 27 \n", - "28 0 False train_loader -1.0 28 \n", - "29 0 False train_loader -1.0 29 \n", - "30 0 False train_loader -1.0 30 \n", - "31 0 False train_loader -1.0 31 \n", - "32 0 False train_loader -1.0 32 \n", - "33 0 False train_loader -1.0 33 \n", - "34 0 False train_loader -1.0 34 \n", - "35 0 False train_loader -1.0 35 \n", - "36 0 False train_loader -1.0 36 \n", - "\n", - " member_rank \\\n", - "sample_id annotation_id \n", - "0 0 0.0 \n", - "1 0 0.0 \n", - "2 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - "3 0 0.0 \n", - "4 0 0.0 \n", - "5 0 0.0 \n", - "6 0 0.0 \n", - "7 0 0.0 \n", - "8 0 0.0 \n", - "9 0 0.0 \n", - "10 0 0.0 \n", - "11 0 0.0 \n", - "12 0 0.0 \n", - "13 0 0.0 \n", - "14 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - "15 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - "16 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - "17 0 0.0 \n", - "18 0 0.0 \n", - "19 0 0.0 \n", - "20 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - " 3 NaN \n", - "21 0 0.0 \n", - "22 0 0.0 \n", - "23 0 0.0 \n", - "24 0 0.0 \n", - "25 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - "26 0 0.0 \n", - "27 0 0.0 \n", - "28 0 0.0 \n", - "29 0 0.0 \n", - "30 0 0.0 \n", - "31 0 0.0 \n", - "32 0 0.0 \n", - "33 0 0.0 \n", - "34 0 0.0 \n", - "35 0 0.0 \n", - "36 0 0.0 \n", - "\n", - " img_path \\\n", - "sample_id annotation_id \n", - "0 0 /content/datasets/brain-tumor/images/train/000... \n", - "1 0 /content/datasets/brain-tumor/images/train/000... \n", - "2 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - "3 0 /content/datasets/brain-tumor/images/train/000... \n", - "4 0 /content/datasets/brain-tumor/images/train/000... \n", - "5 0 /content/datasets/brain-tumor/images/train/000... \n", - "6 0 /content/datasets/brain-tumor/images/train/000... \n", - "7 0 /content/datasets/brain-tumor/images/train/000... \n", - "8 0 /content/datasets/brain-tumor/images/train/000... \n", - "9 0 /content/datasets/brain-tumor/images/train/000... \n", - "10 0 /content/datasets/brain-tumor/images/train/000... \n", - "11 0 /content/datasets/brain-tumor/images/train/000... \n", - "12 0 /content/datasets/brain-tumor/images/train/000... \n", - "13 0 /content/datasets/brain-tumor/images/train/000... \n", - "14 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - "15 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - "16 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - "17 0 /content/datasets/brain-tumor/images/train/000... \n", - "18 0 /content/datasets/brain-tumor/images/train/000... \n", - "19 0 /content/datasets/brain-tumor/images/train/000... \n", - "20 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - " 3 None \n", - "21 0 /content/datasets/brain-tumor/images/train/000... \n", - "22 0 /content/datasets/brain-tumor/images/train/000... \n", - "23 0 /content/datasets/brain-tumor/images/train/000... \n", - "24 0 /content/datasets/brain-tumor/images/train/000... \n", - "25 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - "26 0 /content/datasets/brain-tumor/images/train/000... \n", - "27 0 /content/datasets/brain-tumor/images/train/000... \n", - "28 0 /content/datasets/brain-tumor/images/train/000... \n", - "29 0 /content/datasets/brain-tumor/images/train/000... \n", - "30 0 /content/datasets/brain-tumor/images/train/000... \n", - "31 0 /content/datasets/brain-tumor/images/train/000... \n", - "32 0 /content/datasets/brain-tumor/images/train/000... \n", - "33 0 /content/datasets/brain-tumor/images/train/000... \n", - "34 0 /content/datasets/brain-tumor/images/train/000... \n", - "35 0 /content/datasets/brain-tumor/images/train/000... \n", - "36 0 /content/datasets/brain-tumor/images/train/000... \n", - "\n", - " cls \n", - "sample_id annotation_id \n", - "0 0 [[1.0]] \n", - "1 0 [[1.0]] \n", - "2 0 [[1.0], [1.0]] \n", - " 1 None \n", - " 2 None \n", - "3 0 [[1.0]] \n", - "4 0 [[1.0]] \n", - "5 0 [[1.0]] \n", - "6 0 [[1.0]] \n", - "7 0 [[1.0]] \n", - "8 0 [[1.0]] \n", - "9 0 [[1.0]] \n", - "10 0 [[1.0]] \n", - "11 0 [[1.0]] \n", - "12 0 [[1.0]] \n", - "13 0 [[1.0]] \n", - "14 0 [[1.0], [1.0]] \n", - " 1 None \n", - " 2 None \n", - "15 0 [[1.0], [1.0]] \n", - " 1 None \n", - " 2 None \n", - "16 0 [[1.0], [1.0]] \n", - " 1 None \n", - " 2 None \n", - "17 0 [[1.0]] \n", - "18 0 [[1.0]] \n", - "19 0 [[0.0]] \n", - "20 0 [[0.0], [0.0], [0.0]] \n", - " 1 None \n", - " 2 None \n", - " 3 None \n", - "21 0 [[0.0]] \n", - "22 0 [[0.0]] \n", - "23 0 [[1.0]] \n", - "24 0 [[1.0]] \n", - "25 0 [[0.0], [0.0]] \n", - " 1 None \n", - " 2 None \n", - "26 0 [[0.0]] \n", - "27 0 [[0.0]] \n", - "28 0 [[1.0]] \n", - "29 0 [[1.0]] \n", - "30 0 [[1.0]] \n", - "31 0 [[1.0]] \n", - "32 0 [[1.0]] \n", - "33 0 [[1.0]] \n", - "34 0 [[1.0]] \n", - "35 0 [[1.0]] \n", - "36 0 [[1.0]] " - ] - }, - "metadata": {}, - "output_type": "display_data" - }, - { - "data": { - "text/html": [ - "🔗 Open the full grid in Weights Studio" - ], - "text/plain": [ - "" - ] - }, - "metadata": {}, - "output_type": "display_data" - } - ], + "outputs": [], "source": [ "import pandas as pd\n", "from IPython.display import display, HTML\n", diff --git a/weightslab/examples/Notebooks/Ultralytics/wl-how-to-train-ultralytics-yolo-on-medical-pills-dataset.ipynb b/weightslab/examples/Notebooks/Ultralytics/wl-how-to-train-ultralytics-yolo-on-medical-pills-dataset.ipynb index c0b09ee9..0c6bd306 100644 --- a/weightslab/examples/Notebooks/Ultralytics/wl-how-to-train-ultralytics-yolo-on-medical-pills-dataset.ipynb +++ b/weightslab/examples/Notebooks/Ultralytics/wl-how-to-train-ultralytics-yolo-on-medical-pills-dataset.ipynb @@ -44,7 +44,7 @@ }, { "cell_type": "code", - "execution_count": 6, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/" @@ -52,24 +52,14 @@ "id": "71be394c", "outputId": "c9ef4767-84f5-4bea-e395-9e11766002f2" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Requirement already satisfied: setuptools in /usr/local/lib/python3.12/dist-packages (75.2.0)\n", - "Requirement already satisfied: wheel in /usr/local/lib/python3.12/dist-packages (0.47.0)\n", - "Requirement already satisfied: packaging>=24.0 in /usr/local/lib/python3.12/dist-packages (from wheel) (26.2)\n" - ] - } - ], + "outputs": [], "source": [ "%pip install setuptools wheel" ] }, { "cell_type": "code", - "execution_count": 7, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/", @@ -78,149 +68,7 @@ "id": "jV1oi_PQN0xK", "outputId": "eda8eca1-f4b3-4ffd-eaa6-0f67ce857153" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Collecting git+https://github.com/GrayboxTech/weightslab.git\n", - " Cloning https://github.com/GrayboxTech/weightslab.git (to revision landingcollab) to /tmp/pip-req-build-z4ba7cx9\n", - " Running command git clone --filter=blob:none --quiet https://github.com/GrayboxTech/weightslab.git /tmp/pip-req-build-z4ba7cx9\n", - " Running command git checkout -b landingcollab --track origin/landingcollab\n", - " Switched to a new branch 'landingcollab'\n", - " Branch 'landingcollab' set up to track remote branch 'landingcollab' from 'origin'.\n", - " Resolved https://github.com/GrayboxTech/weightslab.git to commit 3f1cf143bc93e67df43adc71ab831fa71477e0d2\n", - " Installing build dependencies ... \u001b[?25l\u001b[?25hdone\n", - " Getting requirements to build wheel ... \u001b[?25l\u001b[?25hdone\n", - " Preparing metadata (pyproject.toml) ... \u001b[?25l\u001b[?25hdone\n", - "Requirement already satisfied: numpy<3,>=1.24 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (2.0.2)\n", - "Requirement already satisfied: pandas<3,>=2.2.2 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (2.2.2)\n", - "Requirement already satisfied: duckdb<2,>=1.1 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.3.2)\n", - "Requirement already satisfied: PyYAML<7,>=6.0.3 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (6.0.3)\n", - "Requirement already satisfied: dill<0.5,>=0.3.8 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (0.3.8)\n", - "Requirement already satisfied: zstandard<1,>=0.22 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (0.25.0)\n", - "Requirement already satisfied: h5py<4,>=3.10 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (3.16.0)\n", - "Requirement already satisfied: xxhash<4.1,>=3.4 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (3.7.1)\n", - "Requirement already satisfied: tables<4,>=3.9 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (3.10.2)\n", - "Requirement already satisfied: torch<=2.9,>=2.1 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (2.9.0)\n", - "Requirement already satisfied: torchvision<1,>=0.16 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (0.24.0)\n", - "Requirement already satisfied: torchmetrics>=1.9 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.9.0)\n", - "Requirement already satisfied: grpcio<2,>=1.80 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.81.1)\n", - "Requirement already satisfied: protobuf<8,>=5.28.1 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (5.29.6)\n", - "Requirement already satisfied: pydantic<3,>=2.7 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (2.13.4)\n", - "Requirement already satisfied: Pillow<12,>=10 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (11.3.0)\n", - "Requirement already satisfied: graphviz<1,>=0.20 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (0.21)\n", - "Requirement already satisfied: onnx<=1.20,>=1.15 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.20.0)\n", - "Requirement already satisfied: tqdm<5,>=4.66 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (4.67.3)\n", - "Requirement already satisfied: python-dotenv<2,>=1 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.2.2)\n", - "Requirement already satisfied: langchain-core<2,>=0.3 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.4.9)\n", - "Requirement already satisfied: langchain-ollama<2,>=0.2 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.1.0)\n", - "Requirement already satisfied: langchain-openai<2,>=0.2 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.3.5)\n", - "Requirement already satisfied: typing-extensions~=4.12 in /usr/local/lib/python3.12/dist-packages (from grpcio<2,>=1.80->weightslab==1.3.3.dev8) (4.15.0)\n", - "Requirement already satisfied: jsonpatch<2.0.0,>=1.33.0 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (1.33)\n", - "Requirement already satisfied: langchain-protocol>=0.0.17 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (0.0.18)\n", - "Requirement already satisfied: langsmith<1.0.0,>=0.3.45 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (0.9.1)\n", - "Requirement already satisfied: packaging>=23.2.0 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (26.2)\n", - "Requirement already satisfied: tenacity!=8.4.0,<10.0.0,>=8.1.0 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (9.1.4)\n", - "Requirement already satisfied: uuid-utils<1.0,>=0.12.0 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (0.16.2)\n", - "Requirement already satisfied: ollama<1.0.0,>=0.6.1 in /usr/local/lib/python3.12/dist-packages (from langchain-ollama<2,>=0.2->weightslab==1.3.3.dev8) (0.6.2)\n", - "Requirement already satisfied: openai<3.0.0,>=2.45.0 in /usr/local/lib/python3.12/dist-packages (from langchain-openai<2,>=0.2->weightslab==1.3.3.dev8) (2.45.0)\n", - "Requirement already satisfied: tiktoken<1.0.0,>=0.7.0 in /usr/local/lib/python3.12/dist-packages (from langchain-openai<2,>=0.2->weightslab==1.3.3.dev8) (0.13.0)\n", - "Requirement already satisfied: ml_dtypes>=0.5.0 in /usr/local/lib/python3.12/dist-packages (from onnx<=1.20,>=1.15->weightslab==1.3.3.dev8) (0.5.4)\n", - "Requirement already satisfied: python-dateutil>=2.8.2 in /usr/local/lib/python3.12/dist-packages (from pandas<3,>=2.2.2->weightslab==1.3.3.dev8) (2.9.0.post0)\n", - "Requirement already satisfied: pytz>=2020.1 in /usr/local/lib/python3.12/dist-packages (from pandas<3,>=2.2.2->weightslab==1.3.3.dev8) (2025.2)\n", - "Requirement already satisfied: tzdata>=2022.7 in /usr/local/lib/python3.12/dist-packages (from pandas<3,>=2.2.2->weightslab==1.3.3.dev8) (2026.2)\n", - "Requirement already satisfied: annotated-types>=0.6.0 in /usr/local/lib/python3.12/dist-packages (from pydantic<3,>=2.7->weightslab==1.3.3.dev8) (0.7.0)\n", - "Requirement already satisfied: pydantic-core==2.46.4 in /usr/local/lib/python3.12/dist-packages (from pydantic<3,>=2.7->weightslab==1.3.3.dev8) (2.46.4)\n", - "Requirement already satisfied: typing-inspection>=0.4.2 in /usr/local/lib/python3.12/dist-packages (from pydantic<3,>=2.7->weightslab==1.3.3.dev8) (0.4.2)\n", - "Requirement already satisfied: numexpr>=2.6.2 in /usr/local/lib/python3.12/dist-packages (from tables<4,>=3.9->weightslab==1.3.3.dev8) (2.14.1)\n", - "Requirement already satisfied: py-cpuinfo in /usr/local/lib/python3.12/dist-packages (from tables<4,>=3.9->weightslab==1.3.3.dev8) (9.0.0)\n", - "Requirement already satisfied: blosc2>=2.3.0 in /usr/local/lib/python3.12/dist-packages (from tables<4,>=3.9->weightslab==1.3.3.dev8) (4.5.1)\n", - "Requirement already satisfied: filelock in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.29.4)\n", - "Requirement already satisfied: setuptools in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (75.2.0)\n", - "Requirement already satisfied: sympy>=1.13.3 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (1.14.0)\n", - "Requirement already satisfied: networkx>=2.5.1 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.6.1)\n", - "Requirement already satisfied: jinja2 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.1.6)\n", - "Requirement already satisfied: fsspec>=0.8.5 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (2025.3.0)\n", - "Requirement already satisfied: nvidia-cuda-nvrtc-cu12==12.8.93 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.93)\n", - "Requirement already satisfied: nvidia-cuda-runtime-cu12==12.8.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.90)\n", - "Requirement already satisfied: nvidia-cuda-cupti-cu12==12.8.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.90)\n", - "Requirement already satisfied: nvidia-cudnn-cu12==9.10.2.21 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (9.10.2.21)\n", - "Requirement already satisfied: nvidia-cublas-cu12==12.8.4.1 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.4.1)\n", - "Requirement already satisfied: nvidia-cufft-cu12==11.3.3.83 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (11.3.3.83)\n", - "Requirement already satisfied: nvidia-curand-cu12==10.3.9.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (10.3.9.90)\n", - "Requirement already satisfied: nvidia-cusolver-cu12==11.7.3.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (11.7.3.90)\n", - "Requirement already satisfied: nvidia-cusparse-cu12==12.5.8.93 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.5.8.93)\n", - "Requirement already satisfied: nvidia-cusparselt-cu12==0.7.1 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (0.7.1)\n", - "Requirement already satisfied: nvidia-nccl-cu12==2.27.5 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (2.27.5)\n", - "Requirement already satisfied: nvidia-nvshmem-cu12==3.3.20 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.3.20)\n", - "Requirement already satisfied: nvidia-nvtx-cu12==12.8.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.90)\n", - "Requirement already satisfied: nvidia-nvjitlink-cu12==12.8.93 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.93)\n", - "Requirement already satisfied: nvidia-cufile-cu12==1.13.1.3 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (1.13.1.3)\n", - "Requirement already satisfied: triton==3.5.0 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.5.0)\n", - "Requirement already satisfied: lightning-utilities>=0.15.3 in /usr/local/lib/python3.12/dist-packages (from torchmetrics>=1.9->weightslab==1.3.3.dev8) (0.15.3)\n", - "Requirement already satisfied: ndindex in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (1.10.1)\n", - "Requirement already satisfied: msgpack in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (1.2.1)\n", - "Requirement already satisfied: requests in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (2.32.4)\n", - "Requirement already satisfied: rich in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (13.9.4)\n", - "Requirement already satisfied: textual in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (6.2.1)\n", - "Requirement already satisfied: textual-plotext in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (1.0.1)\n", - "Requirement already satisfied: threadpoolctl in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (3.6.0)\n", - "Requirement already satisfied: jsonpointer>=1.9 in /usr/local/lib/python3.12/dist-packages (from jsonpatch<2.0.0,>=1.33.0->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (3.1.1)\n", - "Requirement already satisfied: anyio>=3.5.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (4.14.0)\n", - "Requirement already satisfied: distro>=1.7.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (1.9.0)\n", - "Requirement already satisfied: httpx<1,>=0.23.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (0.28.1)\n", - "Requirement already satisfied: orjson>=3.9.14 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (3.11.9)\n", - "Requirement already satisfied: requests-toolbelt>=1.0.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (1.0.0)\n", - "Requirement already satisfied: sniffio>=1.1 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (1.3.1)\n", - "Requirement already satisfied: websockets>=15.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (15.0.1)\n", - "Requirement already satisfied: jiter<1,>=0.10.0 in /usr/local/lib/python3.12/dist-packages (from openai<3.0.0,>=2.45.0->langchain-openai<2,>=0.2->weightslab==1.3.3.dev8) (0.15.0)\n", - "Requirement already satisfied: six>=1.5 in /usr/local/lib/python3.12/dist-packages (from python-dateutil>=2.8.2->pandas<3,>=2.2.2->weightslab==1.3.3.dev8) (1.17.0)\n", - "Requirement already satisfied: mpmath<1.4,>=1.1.0 in /usr/local/lib/python3.12/dist-packages (from sympy>=1.13.3->torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (1.3.0)\n", - "Requirement already satisfied: regex in /usr/local/lib/python3.12/dist-packages (from tiktoken<1.0.0,>=0.7.0->langchain-openai<2,>=0.2->weightslab==1.3.3.dev8) (2025.11.3)\n", - "Requirement already satisfied: MarkupSafe>=2.0 in /usr/local/lib/python3.12/dist-packages (from jinja2->torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.0.3)\n", - "Requirement already satisfied: idna>=2.8 in /usr/local/lib/python3.12/dist-packages (from anyio>=3.5.0->langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (3.18)\n", - "Requirement already satisfied: certifi in /usr/local/lib/python3.12/dist-packages (from httpx<1,>=0.23.0->langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (2026.6.17)\n", - "Requirement already satisfied: httpcore==1.* in /usr/local/lib/python3.12/dist-packages (from httpx<1,>=0.23.0->langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (1.0.9)\n", - "Requirement already satisfied: h11>=0.16 in /usr/local/lib/python3.12/dist-packages (from httpcore==1.*->httpx<1,>=0.23.0->langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (0.16.0)\n", - "Requirement already satisfied: charset_normalizer<4,>=2 in /usr/local/lib/python3.12/dist-packages (from requests->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (3.4.7)\n", - "Requirement already satisfied: urllib3<3,>=1.21.1 in /usr/local/lib/python3.12/dist-packages (from requests->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (2.5.0)\n", - "Requirement already satisfied: markdown-it-py>=2.2.0 in /usr/local/lib/python3.12/dist-packages (from rich->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (4.2.0)\n", - "Requirement already satisfied: pygments<3.0.0,>=2.13.0 in /usr/local/lib/python3.12/dist-packages (from rich->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (2.20.0)\n", - "Requirement already satisfied: platformdirs<5,>=3.6.0 in /usr/local/lib/python3.12/dist-packages (from textual->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (4.10.0)\n", - "Requirement already satisfied: plotext<6.0.0,>=5.2.8 in /usr/local/lib/python3.12/dist-packages (from textual-plotext->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (5.3.2)\n", - "Requirement already satisfied: mdurl~=0.1 in /usr/local/lib/python3.12/dist-packages (from markdown-it-py>=2.2.0->rich->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (0.1.2)\n", - "Requirement already satisfied: linkify-it-py<3,>=1 in /usr/local/lib/python3.12/dist-packages (from markdown-it-py[linkify,plugins]>=2.1.0->textual->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (2.1.0)\n", - "Requirement already satisfied: mdit-py-plugins>=0.5.0 in /usr/local/lib/python3.12/dist-packages (from markdown-it-py[linkify,plugins]>=2.1.0->textual->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (0.6.1)\n", - "Requirement already satisfied: uc-micro-py in /usr/local/lib/python3.12/dist-packages (from linkify-it-py<3,>=1->markdown-it-py[linkify,plugins]>=2.1.0->textual->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (2.0.0)\n", - "Building wheels for collected packages: weightslab\n", - " Building wheel for weightslab (pyproject.toml) ... \u001b[?25l\u001b[?25hdone\n", - " Created wheel for weightslab: filename=weightslab-1.3.3.dev8-py3-none-any.whl size=2827715 sha256=5c4e7c14bbf95a48b9a8c86b5a8c744af67e1eb395197cbe79d3e443c21cf86e\n", - " Stored in directory: /tmp/pip-ephem-wheel-cache-__zwgf7x/wheels/4d/c7/cd/817c2745790555e7b6c532f5b6055d47567f9416fdfc65879b\n", - "Successfully built weightslab\n", - "Installing collected packages: weightslab\n", - " Attempting uninstall: weightslab\n", - " Found existing installation: weightslab 1.3.3.dev7\n", - " Uninstalling weightslab-1.3.3.dev7:\n", - " Successfully uninstalled weightslab-1.3.3.dev7\n", - "Successfully installed weightslab-1.3.3.dev8\n" - ] - }, - { - "data": { - "application/vnd.colab-display-data+json": { - "id": "b7fce274042740d19087972fe9c01c9b", - "pip_warning": { - "packages": [ - "weightslab" - ] - } - } - }, - "metadata": {}, - "output_type": "display_data" - } - ], + "outputs": [], "source": [ "%pip install weightslab\n", "!pip install --upgrade \"protobuf>=6.31.1\" # Google Collab compat update" @@ -228,7 +76,7 @@ }, { "cell_type": "code", - "execution_count": 1, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/" @@ -236,39 +84,7 @@ "id": "pJhvREbaKs7f", "outputId": "651baa36-9ebb-49da-830b-e6c8a6e2b64f" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Ultralytics 8.4.96 🚀 Python-3.12.13 torch-2.9.0+cu128 CPU (Intel Xeon CPU @ 2.20GHz)\n", - "Setup complete ✅ (2 CPUs, 12.7 GB RAM, 30.5/107.7 GB disk)\n", - "16/07/2026-09:29:29.683 INFO:root:setup_logging: WeightsLab logging initialized - Log file: /tmp/tmp6meuzqdz/weightslab_logs/weightslab_20260716_092929.log\n", - "16/07/2026-09:29:29.685 INFO:weightslab:: WeightsLab package initialized - Log level: INFO, Log to file: True\n", - "16/07/2026-09:29:29.686 INFO:weightslab:: \n", - "╭─────────────────────────────────────────────────────────────────────────────────────────────────────╮\n", - "│ │\n", - "│ \u001b[32m$$\\ $$\\ \u001b[0m $$\\ $$\\ $$\\ \u001b[31m$$\\ \u001b[0m $$\\ │\n", - "│ \u001b[32m$$ | $\\ $$ |\u001b[0m \\__| $$ | $$ | \u001b[31m$$ | \u001b[0m $$ | │\n", - "│ \u001b[32m$$ |$$$\\ $$ |\u001b[0m $$$$$$\\ $$\\ $$$$$$\\ $$$$$$$\\ $$$$$$\\ $$$$$$$\\ \u001b[31m$$ | \u001b[0m $$$$$$\\ $$$$$$$\\ │\n", - "│ \u001b[32m$$ $$ $$\\$$ |\u001b[0m$$ __$$\\ $$ |$$ __$$\\ $$ __$$\\ \\_$$ _| $$ _____|\u001b[31m$$ | \u001b[0m \\____$$\\ $$ __$$\\ │\n", - "│ \u001b[32m$$$$ _$$$$ |\u001b[0m$$$$$$$$ |$$ |$$ / $$ |$$ | $$ | $$ | \\$$$$$$\\ \u001b[31m$$ | \u001b[0m $$$$$$$ |$$ | $$ | │\n", - "│ \u001b[32m$$$ / \\$$$ |\u001b[0m$$ ____|$$ |$$ | $$ |$$ | $$ | $$ |$$\\ \\____$$\\ \u001b[31m$$ | \u001b[0m$$ __$$ |$$ | $$ | │\n", - "│ \u001b[32m$$ / \\$$ |\u001b[0m\\$$$$$$$\\ $$ |\\$$$$$$$ |$$ | $$ | \\$$$$ |$$$$$$$ |\u001b[31m$$$$$$$$\\ \u001b[0m\\$$$$$$$ |$$$$$$$ | │\n", - "│ \u001b[32m\\__/ \\__|\u001b[0m \\_______|\\__| \\____$$ |\\__| \\__| \\____/ \\_______/ \u001b[31m\\________|\u001b[0m \\_______|\\_______/ │\n", - "│ \u001b[32m \u001b[0m $$\\ $$ |\u001b[31m\u001b[0m │\n", - "│ \u001b[32m \u001b[0m \\$$$$$$ |\u001b[31m\u001b[0m │\n", - "│ \u001b[32m \u001b[0m \\______/\u001b[31m\u001b[0m │\n", - "│ │\n", - "│ Inspect - Edit - Evolve Neural Networks │\n", - "│ │\n", - "╰─── By GrayBx, v1.3.3.dev8 ──────────────────────────────────────────────────────────────────────────╯\n", - "\n", - "\n", - "16/07/2026-09:29:29.690 INFO:weightslab.utils.telemetry:ping_import: WeightsLab uses anonymous usage data (package version used, and OS name). Set WL_NO_TELEMETRY=1 to disable.\n" - ] - } - ], + "outputs": [], "source": [ "# Install WeightsLab. Colab already ships torch, torchvision, numpy, scikit-learn\n", "# and Pillow, so nothing extra is needed here.\n", @@ -349,7 +165,6 @@ "import warnings\n", "warnings.filterwarnings(\"ignore\")\n", "\n", - "from weightslab.integrations.ultralytics import WLAwareTrainer\n", "from ultralytics import YOLO\n", "\n", "# Hyperparameters registered with WeightsLab -> live-editable from Weights Studio.\n", @@ -374,7 +189,7 @@ "# Read back the (now live) hyperparameters and hand training to WLAwareTrainer.\n", "model = YOLO(cfg[\"model\"])\n", "results = model.train(\n", - " trainer=WLAwareTrainer,\n", + " trainer=wl.WLAwareTrainer,\n", " data=str(cfg[\"data\"]),\n", " imgsz=cfg[\"image_size\"],\n", " epochs=int(cfg[\"epochs\"]),\n", @@ -404,7 +219,7 @@ }, { "cell_type": "code", - "execution_count": 3, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/", @@ -414,1230 +229,7 @@ "id": "uOJdXGA3Ks7o", "outputId": "26652fa3-6a15-43ce-9d75-034eaba21ccd" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "1241 samples tracked\n" - ] - }, - { - "data": { - "application/vnd.google.colaboratory.intrinsic+json": { - "repr_error": "Out of range float values are not JSON compliant: nan", - "type": "dataframe" - }, - "text/html": [ - "\n", - "
\n", - "
\n", - "\n", - "\n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - "
predictionprediction_rawtargetdiscardedorigintask_typelast_seengroup_idmember_rankimg_pathcls
sample_idannotation_id
00NoneNone[0.2335685, 0.254695, 0.4553995, 0.43075103, 1...Falsetrain_loader-1.000.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
10NoneNone[0.251174, 0.24061048, 0.44366202, 0.4307515, ...Falsetrain_loader-1.010.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
20NoneNoneNoneFalsetrain_loader-1.020.0/content/datasets/brain-tumor/images/train/000...[[1.0], [1.0]]
1NoneNone[0.523474, 0.3978875, 0.634976, 0.48122054, 1....NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.40610352, 0.415493, 0.5234745, 0.522301, 1....NoneNaNNoneNaNNoneNaNNoneNone
30NoneNone[0.403756, 0.3826295, 0.637324, 0.5152585, 1.0...Falsetrain_loader-1.030.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
40NoneNone[0.4436615, 0.3814555, 0.5938965, 0.4518785, 1...Falsetrain_loader-1.040.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
50NoneNone[0.6044605, 0.3990615, 0.67253554, 0.46596247,...Falsetrain_loader-1.050.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
60NoneNone[0.410798, 0.4436625, 0.556338, 0.51173747, 1....Falsetrain_loader-1.060.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
70NoneNone[0.650235, 0.4166665, 0.73826295, 0.48356748, ...Falsetrain_loader-1.070.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
80NoneNone[0.64788747, 0.388498, 0.7582165, 0.49413198, ...Falsetrain_loader-1.080.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
90NoneNone[0.6807515, 0.40140802, 0.75704247, 0.45774597...Falsetrain_loader-1.090.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
100NoneNone[0.431925, 0.1701875, 0.561033, 0.3180745, 1.0...Falsetrain_loader-1.0100.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
110NoneNone[0.42018852, 0.34507048, 0.5434275, 0.4941315,...Falsetrain_loader-1.0110.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
120NoneNone[0.3732395, 0.348592, 0.54812247, 0.46009403, ...Falsetrain_loader-1.0120.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
130NoneNone[0.37793398, 0.3814555, 0.48122, 0.45892054, 1...Falsetrain_loader-1.0130.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
140NoneNoneNoneFalsetrain_loader-1.0140.0/content/datasets/brain-tumor/images/train/000...[[1.0], [1.0]]
1NoneNone[0.43427247, 0.3556335, 0.51173747, 0.43075046...NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.4929575, 0.44718304, 0.5375585, 0.485915, 1...NoneNaNNoneNaNNoneNaNNoneNone
150NoneNoneNoneFalsetrain_loader-1.0150.0/content/datasets/brain-tumor/images/train/000...[[1.0], [1.0]]
1NoneNone[0.629108, 0.39788702, 0.67840403, 0.453051, 1...NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.4530515, 0.4107985, 0.58568054, 0.5105635, ...NoneNaNNoneNaNNoneNaNNoneNone
160NoneNoneNoneFalsetrain_loader-1.0160.0/content/datasets/brain-tumor/images/train/000...[[1.0], [1.0]]
1NoneNone[0.45070451, 0.3967135, 0.5762915, 0.5152585, ...NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.63732404, 0.403756, 0.684272, 0.46009403, 1...NoneNaNNoneNaNNoneNaNNoneNone
170NoneNone[0.43779355, 0.3826295, 0.61032856, 0.5223005,...Falsetrain_loader-1.0170.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
180NoneNone[0.4401405, 0.374413, 0.6068075, 0.529343, 1.0...Falsetrain_loader-1.0180.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
190NoneNone[0.536385, 0.4495305, 0.600939, 0.49999946, 0....Falsetrain_loader-1.0190.0/content/datasets/brain-tumor/images/train/000...[[0.0]]
200NoneNoneNoneFalsetrain_loader-1.0200.0/content/datasets/brain-tumor/images/train/000...[[0.0], [0.0], [0.0]]
1NoneNone[0.52347445, 0.4424885, 0.6150235, 0.5281695, ...NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.456573, 0.515258, 0.504695, 0.562206, 0.0, ...NoneNaNNoneNaNNoneNaNNoneNone
3NoneNone[0.504695, 0.44248852, 0.54342705, 0.48004752,...NoneNaNNoneNaNNoneNaNNoneNone
210NoneNone[0.54107946, 0.456573, 0.6103285, 0.511737, 0....Falsetrain_loader-1.0210.0/content/datasets/brain-tumor/images/train/000...[[0.0]]
220NoneNone[0.3298125, 0.2042255, 0.4225355, 0.2957745, 0...Falsetrain_loader-1.0220.0/content/datasets/brain-tumor/images/train/000...[[0.0]]
230NoneNone[0.517606, 0.18661949, 0.615024, 0.2875585, 1....Falsetrain_loader-1.0230.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
240NoneNone[0.51173747, 0.18779299, 0.6126765, 0.308685, ...Falsetrain_loader-1.0240.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
250NoneNoneNoneFalsetrain_loader-1.0250.0/content/datasets/brain-tumor/images/train/000...[[0.0], [0.0]]
1NoneNone[0.30985948, 0.3626765, 0.3861505, 0.43896753,...NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.41314548, 0.3427225, 0.48122054, 0.42018753...NoneNaNNoneNaNNoneNaNNoneNone
260NoneNone[0.25821602, 0.30046952, 0.46830997, 0.4577465...Falsetrain_loader-1.0260.0/content/datasets/brain-tumor/images/train/000...[[0.0]]
270NoneNone[0.2734745, 0.2828635, 0.44718352, 0.4577465, ...Falsetrain_loader-1.0270.0/content/datasets/brain-tumor/images/train/000...[[0.0]]
280NoneNone[0.5915495, 0.23591551, 0.6455405, 0.2887325, ...Falsetrain_loader-1.0280.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
290NoneNone[0.5903755, 0.231221, 0.6420185, 0.280517, 1.0...Falsetrain_loader-1.0290.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
300NoneNone[0.40493003, 0.2018775, 0.469484, 0.2699525, 1...Falsetrain_loader-1.0300.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
310NoneNone[0.39554, 0.19483551, 0.483568, 0.2887325, 1.0...Falsetrain_loader-1.0310.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
320NoneNone[0.4025825, 0.212441, 0.45539945, 0.269953, 1....Falsetrain_loader-1.0320.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
330NoneNone[0.2265255, 0.2136145, 0.3438965, 0.30751148, ...Falsetrain_loader-1.0330.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
340NoneNone[0.634976, 0.349765, 0.7723, 0.456573, 1.0, 1.0]Falsetrain_loader-1.0340.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
350NoneNone[0.6079815, 0.320423, 0.7969485, 0.489437, 1.0...Falsetrain_loader-1.0350.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
360NoneNone[0.6185445, 0.30164298, 0.7758215, 0.494131, 1...Falsetrain_loader-1.0360.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
\n", - "
\n", - "
\n", - "\n", - "
\n", - " \n", - "\n", - " \n", - "\n", - " \n", - "
\n", - "\n", - "\n", - "
\n", - "
\n" - ], - "text/plain": [ - " prediction prediction_raw \\\n", - "sample_id annotation_id \n", - "0 0 None None \n", - "1 0 None None \n", - "2 0 None None \n", - " 1 None None \n", - " 2 None None \n", - "3 0 None None \n", - "4 0 None None \n", - "5 0 None None \n", - "6 0 None None \n", - "7 0 None None \n", - "8 0 None None \n", - "9 0 None None \n", - "10 0 None None \n", - "11 0 None None \n", - "12 0 None None \n", - "13 0 None None \n", - "14 0 None None \n", - " 1 None None \n", - " 2 None None \n", - "15 0 None None \n", - " 1 None None \n", - " 2 None None \n", - "16 0 None None \n", - " 1 None None \n", - " 2 None None \n", - "17 0 None None \n", - "18 0 None None \n", - "19 0 None None \n", - "20 0 None None \n", - " 1 None None \n", - " 2 None None \n", - " 3 None None \n", - "21 0 None None \n", - "22 0 None None \n", - "23 0 None None \n", - "24 0 None None \n", - "25 0 None None \n", - " 1 None None \n", - " 2 None None \n", - "26 0 None None \n", - "27 0 None None \n", - "28 0 None None \n", - "29 0 None None \n", - "30 0 None None \n", - "31 0 None None \n", - "32 0 None None \n", - "33 0 None None \n", - "34 0 None None \n", - "35 0 None None \n", - "36 0 None None \n", - "\n", - " target \\\n", - "sample_id annotation_id \n", - "0 0 [0.2335685, 0.254695, 0.4553995, 0.43075103, 1... \n", - "1 0 [0.251174, 0.24061048, 0.44366202, 0.4307515, ... \n", - "2 0 None \n", - " 1 [0.523474, 0.3978875, 0.634976, 0.48122054, 1.... \n", - " 2 [0.40610352, 0.415493, 0.5234745, 0.522301, 1.... \n", - "3 0 [0.403756, 0.3826295, 0.637324, 0.5152585, 1.0... \n", - "4 0 [0.4436615, 0.3814555, 0.5938965, 0.4518785, 1... \n", - "5 0 [0.6044605, 0.3990615, 0.67253554, 0.46596247,... \n", - "6 0 [0.410798, 0.4436625, 0.556338, 0.51173747, 1.... \n", - "7 0 [0.650235, 0.4166665, 0.73826295, 0.48356748, ... \n", - "8 0 [0.64788747, 0.388498, 0.7582165, 0.49413198, ... \n", - "9 0 [0.6807515, 0.40140802, 0.75704247, 0.45774597... \n", - "10 0 [0.431925, 0.1701875, 0.561033, 0.3180745, 1.0... \n", - "11 0 [0.42018852, 0.34507048, 0.5434275, 0.4941315,... \n", - "12 0 [0.3732395, 0.348592, 0.54812247, 0.46009403, ... \n", - "13 0 [0.37793398, 0.3814555, 0.48122, 0.45892054, 1... \n", - "14 0 None \n", - " 1 [0.43427247, 0.3556335, 0.51173747, 0.43075046... \n", - " 2 [0.4929575, 0.44718304, 0.5375585, 0.485915, 1... \n", - "15 0 None \n", - " 1 [0.629108, 0.39788702, 0.67840403, 0.453051, 1... \n", - " 2 [0.4530515, 0.4107985, 0.58568054, 0.5105635, ... \n", - "16 0 None \n", - " 1 [0.45070451, 0.3967135, 0.5762915, 0.5152585, ... \n", - " 2 [0.63732404, 0.403756, 0.684272, 0.46009403, 1... \n", - "17 0 [0.43779355, 0.3826295, 0.61032856, 0.5223005,... \n", - "18 0 [0.4401405, 0.374413, 0.6068075, 0.529343, 1.0... \n", - "19 0 [0.536385, 0.4495305, 0.600939, 0.49999946, 0.... \n", - "20 0 None \n", - " 1 [0.52347445, 0.4424885, 0.6150235, 0.5281695, ... \n", - " 2 [0.456573, 0.515258, 0.504695, 0.562206, 0.0, ... \n", - " 3 [0.504695, 0.44248852, 0.54342705, 0.48004752,... \n", - "21 0 [0.54107946, 0.456573, 0.6103285, 0.511737, 0.... \n", - "22 0 [0.3298125, 0.2042255, 0.4225355, 0.2957745, 0... \n", - "23 0 [0.517606, 0.18661949, 0.615024, 0.2875585, 1.... \n", - "24 0 [0.51173747, 0.18779299, 0.6126765, 0.308685, ... \n", - "25 0 None \n", - " 1 [0.30985948, 0.3626765, 0.3861505, 0.43896753,... \n", - " 2 [0.41314548, 0.3427225, 0.48122054, 0.42018753... \n", - "26 0 [0.25821602, 0.30046952, 0.46830997, 0.4577465... \n", - "27 0 [0.2734745, 0.2828635, 0.44718352, 0.4577465, ... \n", - "28 0 [0.5915495, 0.23591551, 0.6455405, 0.2887325, ... \n", - "29 0 [0.5903755, 0.231221, 0.6420185, 0.280517, 1.0... \n", - "30 0 [0.40493003, 0.2018775, 0.469484, 0.2699525, 1... \n", - "31 0 [0.39554, 0.19483551, 0.483568, 0.2887325, 1.0... \n", - "32 0 [0.4025825, 0.212441, 0.45539945, 0.269953, 1.... \n", - "33 0 [0.2265255, 0.2136145, 0.3438965, 0.30751148, ... \n", - "34 0 [0.634976, 0.349765, 0.7723, 0.456573, 1.0, 1.0] \n", - "35 0 [0.6079815, 0.320423, 0.7969485, 0.489437, 1.0... \n", - "36 0 [0.6185445, 0.30164298, 0.7758215, 0.494131, 1... \n", - "\n", - " discarded origin task_type last_seen group_id \\\n", - "sample_id annotation_id \n", - "0 0 False train_loader -1.0 0 \n", - "1 0 False train_loader -1.0 1 \n", - "2 0 False train_loader -1.0 2 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - "3 0 False train_loader -1.0 3 \n", - "4 0 False train_loader -1.0 4 \n", - "5 0 False train_loader -1.0 5 \n", - "6 0 False train_loader -1.0 6 \n", - "7 0 False train_loader -1.0 7 \n", - "8 0 False train_loader -1.0 8 \n", - "9 0 False train_loader -1.0 9 \n", - "10 0 False train_loader -1.0 10 \n", - "11 0 False train_loader -1.0 11 \n", - "12 0 False train_loader -1.0 12 \n", - "13 0 False train_loader -1.0 13 \n", - "14 0 False train_loader -1.0 14 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - "15 0 False train_loader -1.0 15 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - "16 0 False train_loader -1.0 16 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - "17 0 False train_loader -1.0 17 \n", - "18 0 False train_loader -1.0 18 \n", - "19 0 False train_loader -1.0 19 \n", - "20 0 False train_loader -1.0 20 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - " 3 None NaN None NaN None \n", - "21 0 False train_loader -1.0 21 \n", - "22 0 False train_loader -1.0 22 \n", - "23 0 False train_loader -1.0 23 \n", - "24 0 False train_loader -1.0 24 \n", - "25 0 False train_loader -1.0 25 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - "26 0 False train_loader -1.0 26 \n", - "27 0 False train_loader -1.0 27 \n", - "28 0 False train_loader -1.0 28 \n", - "29 0 False train_loader -1.0 29 \n", - "30 0 False train_loader -1.0 30 \n", - "31 0 False train_loader -1.0 31 \n", - "32 0 False train_loader -1.0 32 \n", - "33 0 False train_loader -1.0 33 \n", - "34 0 False train_loader -1.0 34 \n", - "35 0 False train_loader -1.0 35 \n", - "36 0 False train_loader -1.0 36 \n", - "\n", - " member_rank \\\n", - "sample_id annotation_id \n", - "0 0 0.0 \n", - "1 0 0.0 \n", - "2 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - "3 0 0.0 \n", - "4 0 0.0 \n", - "5 0 0.0 \n", - "6 0 0.0 \n", - "7 0 0.0 \n", - "8 0 0.0 \n", - "9 0 0.0 \n", - "10 0 0.0 \n", - "11 0 0.0 \n", - "12 0 0.0 \n", - "13 0 0.0 \n", - "14 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - "15 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - "16 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - "17 0 0.0 \n", - "18 0 0.0 \n", - "19 0 0.0 \n", - "20 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - " 3 NaN \n", - "21 0 0.0 \n", - "22 0 0.0 \n", - "23 0 0.0 \n", - "24 0 0.0 \n", - "25 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - "26 0 0.0 \n", - "27 0 0.0 \n", - "28 0 0.0 \n", - "29 0 0.0 \n", - "30 0 0.0 \n", - "31 0 0.0 \n", - "32 0 0.0 \n", - "33 0 0.0 \n", - "34 0 0.0 \n", - "35 0 0.0 \n", - "36 0 0.0 \n", - "\n", - " img_path \\\n", - "sample_id annotation_id \n", - "0 0 /content/datasets/brain-tumor/images/train/000... \n", - "1 0 /content/datasets/brain-tumor/images/train/000... \n", - "2 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - "3 0 /content/datasets/brain-tumor/images/train/000... \n", - "4 0 /content/datasets/brain-tumor/images/train/000... \n", - "5 0 /content/datasets/brain-tumor/images/train/000... \n", - "6 0 /content/datasets/brain-tumor/images/train/000... \n", - "7 0 /content/datasets/brain-tumor/images/train/000... \n", - "8 0 /content/datasets/brain-tumor/images/train/000... \n", - "9 0 /content/datasets/brain-tumor/images/train/000... \n", - "10 0 /content/datasets/brain-tumor/images/train/000... \n", - "11 0 /content/datasets/brain-tumor/images/train/000... \n", - "12 0 /content/datasets/brain-tumor/images/train/000... \n", - "13 0 /content/datasets/brain-tumor/images/train/000... \n", - "14 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - "15 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - "16 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - "17 0 /content/datasets/brain-tumor/images/train/000... \n", - "18 0 /content/datasets/brain-tumor/images/train/000... \n", - "19 0 /content/datasets/brain-tumor/images/train/000... \n", - "20 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - " 3 None \n", - "21 0 /content/datasets/brain-tumor/images/train/000... \n", - "22 0 /content/datasets/brain-tumor/images/train/000... \n", - "23 0 /content/datasets/brain-tumor/images/train/000... \n", - "24 0 /content/datasets/brain-tumor/images/train/000... \n", - "25 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - "26 0 /content/datasets/brain-tumor/images/train/000... \n", - "27 0 /content/datasets/brain-tumor/images/train/000... \n", - "28 0 /content/datasets/brain-tumor/images/train/000... \n", - "29 0 /content/datasets/brain-tumor/images/train/000... \n", - "30 0 /content/datasets/brain-tumor/images/train/000... \n", - "31 0 /content/datasets/brain-tumor/images/train/000... \n", - "32 0 /content/datasets/brain-tumor/images/train/000... \n", - "33 0 /content/datasets/brain-tumor/images/train/000... \n", - "34 0 /content/datasets/brain-tumor/images/train/000... \n", - "35 0 /content/datasets/brain-tumor/images/train/000... \n", - "36 0 /content/datasets/brain-tumor/images/train/000... \n", - "\n", - " cls \n", - "sample_id annotation_id \n", - "0 0 [[1.0]] \n", - "1 0 [[1.0]] \n", - "2 0 [[1.0], [1.0]] \n", - " 1 None \n", - " 2 None \n", - "3 0 [[1.0]] \n", - "4 0 [[1.0]] \n", - "5 0 [[1.0]] \n", - "6 0 [[1.0]] \n", - "7 0 [[1.0]] \n", - "8 0 [[1.0]] \n", - "9 0 [[1.0]] \n", - "10 0 [[1.0]] \n", - "11 0 [[1.0]] \n", - "12 0 [[1.0]] \n", - "13 0 [[1.0]] \n", - "14 0 [[1.0], [1.0]] \n", - " 1 None \n", - " 2 None \n", - "15 0 [[1.0], [1.0]] \n", - " 1 None \n", - " 2 None \n", - "16 0 [[1.0], [1.0]] \n", - " 1 None \n", - " 2 None \n", - "17 0 [[1.0]] \n", - "18 0 [[1.0]] \n", - "19 0 [[0.0]] \n", - "20 0 [[0.0], [0.0], [0.0]] \n", - " 1 None \n", - " 2 None \n", - " 3 None \n", - "21 0 [[0.0]] \n", - "22 0 [[0.0]] \n", - "23 0 [[1.0]] \n", - "24 0 [[1.0]] \n", - "25 0 [[0.0], [0.0]] \n", - " 1 None \n", - " 2 None \n", - "26 0 [[0.0]] \n", - "27 0 [[0.0]] \n", - "28 0 [[1.0]] \n", - "29 0 [[1.0]] \n", - "30 0 [[1.0]] \n", - "31 0 [[1.0]] \n", - "32 0 [[1.0]] \n", - "33 0 [[1.0]] \n", - "34 0 [[1.0]] \n", - "35 0 [[1.0]] \n", - "36 0 [[1.0]] " - ] - }, - "metadata": {}, - "output_type": "display_data" - }, - { - "data": { - "text/html": [ - "🔗 Open the full grid in Weights Studio" - ], - "text/plain": [ - "" - ] - }, - "metadata": {}, - "output_type": "display_data" - } - ], + "outputs": [], "source": [ "import pandas as pd\n", "from IPython.display import display, HTML\n", diff --git a/weightslab/examples/Notebooks/Ultralytics/wl-how-to-train-ultralytics-yolo-on-package-segmentation-dataset.ipynb b/weightslab/examples/Notebooks/Ultralytics/wl-how-to-train-ultralytics-yolo-on-package-segmentation-dataset.ipynb index b6aa99d7..28349fca 100644 --- a/weightslab/examples/Notebooks/Ultralytics/wl-how-to-train-ultralytics-yolo-on-package-segmentation-dataset.ipynb +++ b/weightslab/examples/Notebooks/Ultralytics/wl-how-to-train-ultralytics-yolo-on-package-segmentation-dataset.ipynb @@ -44,7 +44,7 @@ }, { "cell_type": "code", - "execution_count": 6, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/" @@ -52,24 +52,14 @@ "id": "71be394c", "outputId": "c9ef4767-84f5-4bea-e395-9e11766002f2" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Requirement already satisfied: setuptools in /usr/local/lib/python3.12/dist-packages (75.2.0)\n", - "Requirement already satisfied: wheel in /usr/local/lib/python3.12/dist-packages (0.47.0)\n", - "Requirement already satisfied: packaging>=24.0 in /usr/local/lib/python3.12/dist-packages (from wheel) (26.2)\n" - ] - } - ], + "outputs": [], "source": [ "%pip install setuptools wheel" ] }, { "cell_type": "code", - "execution_count": 7, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/", @@ -78,149 +68,7 @@ "id": "jV1oi_PQN0xK", "outputId": "eda8eca1-f4b3-4ffd-eaa6-0f67ce857153" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Collecting git+https://github.com/GrayboxTech/weightslab.git\n", - " Cloning https://github.com/GrayboxTech/weightslab.git (to revision landingcollab) to /tmp/pip-req-build-z4ba7cx9\n", - " Running command git clone --filter=blob:none --quiet https://github.com/GrayboxTech/weightslab.git /tmp/pip-req-build-z4ba7cx9\n", - " Running command git checkout -b landingcollab --track origin/landingcollab\n", - " Switched to a new branch 'landingcollab'\n", - " Branch 'landingcollab' set up to track remote branch 'landingcollab' from 'origin'.\n", - " Resolved https://github.com/GrayboxTech/weightslab.git to commit 3f1cf143bc93e67df43adc71ab831fa71477e0d2\n", - " Installing build dependencies ... \u001b[?25l\u001b[?25hdone\n", - " Getting requirements to build wheel ... \u001b[?25l\u001b[?25hdone\n", - " Preparing metadata (pyproject.toml) ... \u001b[?25l\u001b[?25hdone\n", - "Requirement already satisfied: numpy<3,>=1.24 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (2.0.2)\n", - "Requirement already satisfied: pandas<3,>=2.2.2 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (2.2.2)\n", - "Requirement already satisfied: duckdb<2,>=1.1 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.3.2)\n", - "Requirement already satisfied: PyYAML<7,>=6.0.3 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (6.0.3)\n", - "Requirement already satisfied: dill<0.5,>=0.3.8 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (0.3.8)\n", - "Requirement already satisfied: zstandard<1,>=0.22 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (0.25.0)\n", - "Requirement already satisfied: h5py<4,>=3.10 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (3.16.0)\n", - "Requirement already satisfied: xxhash<4.1,>=3.4 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (3.7.1)\n", - "Requirement already satisfied: tables<4,>=3.9 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (3.10.2)\n", - "Requirement already satisfied: torch<=2.9,>=2.1 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (2.9.0)\n", - "Requirement already satisfied: torchvision<1,>=0.16 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (0.24.0)\n", - "Requirement already satisfied: torchmetrics>=1.9 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.9.0)\n", - "Requirement already satisfied: grpcio<2,>=1.80 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.81.1)\n", - "Requirement already satisfied: protobuf<8,>=5.28.1 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (5.29.6)\n", - "Requirement already satisfied: pydantic<3,>=2.7 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (2.13.4)\n", - "Requirement already satisfied: Pillow<12,>=10 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (11.3.0)\n", - "Requirement already satisfied: graphviz<1,>=0.20 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (0.21)\n", - "Requirement already satisfied: onnx<=1.20,>=1.15 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.20.0)\n", - "Requirement already satisfied: tqdm<5,>=4.66 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (4.67.3)\n", - "Requirement already satisfied: python-dotenv<2,>=1 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.2.2)\n", - "Requirement already satisfied: langchain-core<2,>=0.3 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.4.9)\n", - "Requirement already satisfied: langchain-ollama<2,>=0.2 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.1.0)\n", - "Requirement already satisfied: langchain-openai<2,>=0.2 in /usr/local/lib/python3.12/dist-packages (from weightslab==1.3.3.dev8) (1.3.5)\n", - "Requirement already satisfied: typing-extensions~=4.12 in /usr/local/lib/python3.12/dist-packages (from grpcio<2,>=1.80->weightslab==1.3.3.dev8) (4.15.0)\n", - "Requirement already satisfied: jsonpatch<2.0.0,>=1.33.0 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (1.33)\n", - "Requirement already satisfied: langchain-protocol>=0.0.17 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (0.0.18)\n", - "Requirement already satisfied: langsmith<1.0.0,>=0.3.45 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (0.9.1)\n", - "Requirement already satisfied: packaging>=23.2.0 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (26.2)\n", - "Requirement already satisfied: tenacity!=8.4.0,<10.0.0,>=8.1.0 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (9.1.4)\n", - "Requirement already satisfied: uuid-utils<1.0,>=0.12.0 in /usr/local/lib/python3.12/dist-packages (from langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (0.16.2)\n", - "Requirement already satisfied: ollama<1.0.0,>=0.6.1 in /usr/local/lib/python3.12/dist-packages (from langchain-ollama<2,>=0.2->weightslab==1.3.3.dev8) (0.6.2)\n", - "Requirement already satisfied: openai<3.0.0,>=2.45.0 in /usr/local/lib/python3.12/dist-packages (from langchain-openai<2,>=0.2->weightslab==1.3.3.dev8) (2.45.0)\n", - "Requirement already satisfied: tiktoken<1.0.0,>=0.7.0 in /usr/local/lib/python3.12/dist-packages (from langchain-openai<2,>=0.2->weightslab==1.3.3.dev8) (0.13.0)\n", - "Requirement already satisfied: ml_dtypes>=0.5.0 in /usr/local/lib/python3.12/dist-packages (from onnx<=1.20,>=1.15->weightslab==1.3.3.dev8) (0.5.4)\n", - "Requirement already satisfied: python-dateutil>=2.8.2 in /usr/local/lib/python3.12/dist-packages (from pandas<3,>=2.2.2->weightslab==1.3.3.dev8) (2.9.0.post0)\n", - "Requirement already satisfied: pytz>=2020.1 in /usr/local/lib/python3.12/dist-packages (from pandas<3,>=2.2.2->weightslab==1.3.3.dev8) (2025.2)\n", - "Requirement already satisfied: tzdata>=2022.7 in /usr/local/lib/python3.12/dist-packages (from pandas<3,>=2.2.2->weightslab==1.3.3.dev8) (2026.2)\n", - "Requirement already satisfied: annotated-types>=0.6.0 in /usr/local/lib/python3.12/dist-packages (from pydantic<3,>=2.7->weightslab==1.3.3.dev8) (0.7.0)\n", - "Requirement already satisfied: pydantic-core==2.46.4 in /usr/local/lib/python3.12/dist-packages (from pydantic<3,>=2.7->weightslab==1.3.3.dev8) (2.46.4)\n", - "Requirement already satisfied: typing-inspection>=0.4.2 in /usr/local/lib/python3.12/dist-packages (from pydantic<3,>=2.7->weightslab==1.3.3.dev8) (0.4.2)\n", - "Requirement already satisfied: numexpr>=2.6.2 in /usr/local/lib/python3.12/dist-packages (from tables<4,>=3.9->weightslab==1.3.3.dev8) (2.14.1)\n", - "Requirement already satisfied: py-cpuinfo in /usr/local/lib/python3.12/dist-packages (from tables<4,>=3.9->weightslab==1.3.3.dev8) (9.0.0)\n", - "Requirement already satisfied: blosc2>=2.3.0 in /usr/local/lib/python3.12/dist-packages (from tables<4,>=3.9->weightslab==1.3.3.dev8) (4.5.1)\n", - "Requirement already satisfied: filelock in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.29.4)\n", - "Requirement already satisfied: setuptools in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (75.2.0)\n", - "Requirement already satisfied: sympy>=1.13.3 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (1.14.0)\n", - "Requirement already satisfied: networkx>=2.5.1 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.6.1)\n", - "Requirement already satisfied: jinja2 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.1.6)\n", - "Requirement already satisfied: fsspec>=0.8.5 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (2025.3.0)\n", - "Requirement already satisfied: nvidia-cuda-nvrtc-cu12==12.8.93 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.93)\n", - "Requirement already satisfied: nvidia-cuda-runtime-cu12==12.8.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.90)\n", - "Requirement already satisfied: nvidia-cuda-cupti-cu12==12.8.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.90)\n", - "Requirement already satisfied: nvidia-cudnn-cu12==9.10.2.21 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (9.10.2.21)\n", - "Requirement already satisfied: nvidia-cublas-cu12==12.8.4.1 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.4.1)\n", - "Requirement already satisfied: nvidia-cufft-cu12==11.3.3.83 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (11.3.3.83)\n", - "Requirement already satisfied: nvidia-curand-cu12==10.3.9.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (10.3.9.90)\n", - "Requirement already satisfied: nvidia-cusolver-cu12==11.7.3.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (11.7.3.90)\n", - "Requirement already satisfied: nvidia-cusparse-cu12==12.5.8.93 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.5.8.93)\n", - "Requirement already satisfied: nvidia-cusparselt-cu12==0.7.1 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (0.7.1)\n", - "Requirement already satisfied: nvidia-nccl-cu12==2.27.5 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (2.27.5)\n", - "Requirement already satisfied: nvidia-nvshmem-cu12==3.3.20 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.3.20)\n", - "Requirement already satisfied: nvidia-nvtx-cu12==12.8.90 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.90)\n", - "Requirement already satisfied: nvidia-nvjitlink-cu12==12.8.93 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (12.8.93)\n", - "Requirement already satisfied: nvidia-cufile-cu12==1.13.1.3 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (1.13.1.3)\n", - "Requirement already satisfied: triton==3.5.0 in /usr/local/lib/python3.12/dist-packages (from torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.5.0)\n", - "Requirement already satisfied: lightning-utilities>=0.15.3 in /usr/local/lib/python3.12/dist-packages (from torchmetrics>=1.9->weightslab==1.3.3.dev8) (0.15.3)\n", - "Requirement already satisfied: ndindex in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (1.10.1)\n", - "Requirement already satisfied: msgpack in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (1.2.1)\n", - "Requirement already satisfied: requests in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (2.32.4)\n", - "Requirement already satisfied: rich in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (13.9.4)\n", - "Requirement already satisfied: textual in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (6.2.1)\n", - "Requirement already satisfied: textual-plotext in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (1.0.1)\n", - "Requirement already satisfied: threadpoolctl in /usr/local/lib/python3.12/dist-packages (from blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (3.6.0)\n", - "Requirement already satisfied: jsonpointer>=1.9 in /usr/local/lib/python3.12/dist-packages (from jsonpatch<2.0.0,>=1.33.0->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (3.1.1)\n", - "Requirement already satisfied: anyio>=3.5.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (4.14.0)\n", - "Requirement already satisfied: distro>=1.7.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (1.9.0)\n", - "Requirement already satisfied: httpx<1,>=0.23.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (0.28.1)\n", - "Requirement already satisfied: orjson>=3.9.14 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (3.11.9)\n", - "Requirement already satisfied: requests-toolbelt>=1.0.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (1.0.0)\n", - "Requirement already satisfied: sniffio>=1.1 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (1.3.1)\n", - "Requirement already satisfied: websockets>=15.0 in /usr/local/lib/python3.12/dist-packages (from langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (15.0.1)\n", - "Requirement already satisfied: jiter<1,>=0.10.0 in /usr/local/lib/python3.12/dist-packages (from openai<3.0.0,>=2.45.0->langchain-openai<2,>=0.2->weightslab==1.3.3.dev8) (0.15.0)\n", - "Requirement already satisfied: six>=1.5 in /usr/local/lib/python3.12/dist-packages (from python-dateutil>=2.8.2->pandas<3,>=2.2.2->weightslab==1.3.3.dev8) (1.17.0)\n", - "Requirement already satisfied: mpmath<1.4,>=1.1.0 in /usr/local/lib/python3.12/dist-packages (from sympy>=1.13.3->torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (1.3.0)\n", - "Requirement already satisfied: regex in /usr/local/lib/python3.12/dist-packages (from tiktoken<1.0.0,>=0.7.0->langchain-openai<2,>=0.2->weightslab==1.3.3.dev8) (2025.11.3)\n", - "Requirement already satisfied: MarkupSafe>=2.0 in /usr/local/lib/python3.12/dist-packages (from jinja2->torch<=2.9,>=2.1->weightslab==1.3.3.dev8) (3.0.3)\n", - "Requirement already satisfied: idna>=2.8 in /usr/local/lib/python3.12/dist-packages (from anyio>=3.5.0->langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (3.18)\n", - "Requirement already satisfied: certifi in /usr/local/lib/python3.12/dist-packages (from httpx<1,>=0.23.0->langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (2026.6.17)\n", - "Requirement already satisfied: httpcore==1.* in /usr/local/lib/python3.12/dist-packages (from httpx<1,>=0.23.0->langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (1.0.9)\n", - "Requirement already satisfied: h11>=0.16 in /usr/local/lib/python3.12/dist-packages (from httpcore==1.*->httpx<1,>=0.23.0->langsmith<1.0.0,>=0.3.45->langchain-core<2,>=0.3->weightslab==1.3.3.dev8) (0.16.0)\n", - "Requirement already satisfied: charset_normalizer<4,>=2 in /usr/local/lib/python3.12/dist-packages (from requests->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (3.4.7)\n", - "Requirement already satisfied: urllib3<3,>=1.21.1 in /usr/local/lib/python3.12/dist-packages (from requests->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (2.5.0)\n", - "Requirement already satisfied: markdown-it-py>=2.2.0 in /usr/local/lib/python3.12/dist-packages (from rich->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (4.2.0)\n", - "Requirement already satisfied: pygments<3.0.0,>=2.13.0 in /usr/local/lib/python3.12/dist-packages (from rich->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (2.20.0)\n", - "Requirement already satisfied: platformdirs<5,>=3.6.0 in /usr/local/lib/python3.12/dist-packages (from textual->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (4.10.0)\n", - "Requirement already satisfied: plotext<6.0.0,>=5.2.8 in /usr/local/lib/python3.12/dist-packages (from textual-plotext->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (5.3.2)\n", - "Requirement already satisfied: mdurl~=0.1 in /usr/local/lib/python3.12/dist-packages (from markdown-it-py>=2.2.0->rich->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (0.1.2)\n", - "Requirement already satisfied: linkify-it-py<3,>=1 in /usr/local/lib/python3.12/dist-packages (from markdown-it-py[linkify,plugins]>=2.1.0->textual->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (2.1.0)\n", - "Requirement already satisfied: mdit-py-plugins>=0.5.0 in /usr/local/lib/python3.12/dist-packages (from markdown-it-py[linkify,plugins]>=2.1.0->textual->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (0.6.1)\n", - "Requirement already satisfied: uc-micro-py in /usr/local/lib/python3.12/dist-packages (from linkify-it-py<3,>=1->markdown-it-py[linkify,plugins]>=2.1.0->textual->blosc2>=2.3.0->tables<4,>=3.9->weightslab==1.3.3.dev8) (2.0.0)\n", - "Building wheels for collected packages: weightslab\n", - " Building wheel for weightslab (pyproject.toml) ... \u001b[?25l\u001b[?25hdone\n", - " Created wheel for weightslab: filename=weightslab-1.3.3.dev8-py3-none-any.whl size=2827715 sha256=5c4e7c14bbf95a48b9a8c86b5a8c744af67e1eb395197cbe79d3e443c21cf86e\n", - " Stored in directory: /tmp/pip-ephem-wheel-cache-__zwgf7x/wheels/4d/c7/cd/817c2745790555e7b6c532f5b6055d47567f9416fdfc65879b\n", - "Successfully built weightslab\n", - "Installing collected packages: weightslab\n", - " Attempting uninstall: weightslab\n", - " Found existing installation: weightslab 1.3.3.dev7\n", - " Uninstalling weightslab-1.3.3.dev7:\n", - " Successfully uninstalled weightslab-1.3.3.dev7\n", - "Successfully installed weightslab-1.3.3.dev8\n" - ] - }, - { - "data": { - "application/vnd.colab-display-data+json": { - "id": "b7fce274042740d19087972fe9c01c9b", - "pip_warning": { - "packages": [ - "weightslab" - ] - } - } - }, - "metadata": {}, - "output_type": "display_data" - } - ], + "outputs": [], "source": [ "%pip install weightslab\n", "!pip install --upgrade \"protobuf>=6.31.1\" # Google Collab compat update" @@ -228,7 +76,7 @@ }, { "cell_type": "code", - "execution_count": 1, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/" @@ -236,39 +84,7 @@ "id": "pJhvREbaKs7f", "outputId": "651baa36-9ebb-49da-830b-e6c8a6e2b64f" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Ultralytics 8.4.96 🚀 Python-3.12.13 torch-2.9.0+cu128 CPU (Intel Xeon CPU @ 2.20GHz)\n", - "Setup complete ✅ (2 CPUs, 12.7 GB RAM, 30.5/107.7 GB disk)\n", - "16/07/2026-09:29:29.683 INFO:root:setup_logging: WeightsLab logging initialized - Log file: /tmp/tmp6meuzqdz/weightslab_logs/weightslab_20260716_092929.log\n", - "16/07/2026-09:29:29.685 INFO:weightslab:: WeightsLab package initialized - Log level: INFO, Log to file: True\n", - "16/07/2026-09:29:29.686 INFO:weightslab:: \n", - "╭─────────────────────────────────────────────────────────────────────────────────────────────────────╮\n", - "│ │\n", - "│ \u001b[32m$$\\ $$\\ \u001b[0m $$\\ $$\\ $$\\ \u001b[31m$$\\ \u001b[0m $$\\ │\n", - "│ \u001b[32m$$ | $\\ $$ |\u001b[0m \\__| $$ | $$ | \u001b[31m$$ | \u001b[0m $$ | │\n", - "│ \u001b[32m$$ |$$$\\ $$ |\u001b[0m $$$$$$\\ $$\\ $$$$$$\\ $$$$$$$\\ $$$$$$\\ $$$$$$$\\ \u001b[31m$$ | \u001b[0m $$$$$$\\ $$$$$$$\\ │\n", - "│ \u001b[32m$$ $$ $$\\$$ |\u001b[0m$$ __$$\\ $$ |$$ __$$\\ $$ __$$\\ \\_$$ _| $$ _____|\u001b[31m$$ | \u001b[0m \\____$$\\ $$ __$$\\ │\n", - "│ \u001b[32m$$$$ _$$$$ |\u001b[0m$$$$$$$$ |$$ |$$ / $$ |$$ | $$ | $$ | \\$$$$$$\\ \u001b[31m$$ | \u001b[0m $$$$$$$ |$$ | $$ | │\n", - "│ \u001b[32m$$$ / \\$$$ |\u001b[0m$$ ____|$$ |$$ | $$ |$$ | $$ | $$ |$$\\ \\____$$\\ \u001b[31m$$ | \u001b[0m$$ __$$ |$$ | $$ | │\n", - "│ \u001b[32m$$ / \\$$ |\u001b[0m\\$$$$$$$\\ $$ |\\$$$$$$$ |$$ | $$ | \\$$$$ |$$$$$$$ |\u001b[31m$$$$$$$$\\ \u001b[0m\\$$$$$$$ |$$$$$$$ | │\n", - "│ \u001b[32m\\__/ \\__|\u001b[0m \\_______|\\__| \\____$$ |\\__| \\__| \\____/ \\_______/ \u001b[31m\\________|\u001b[0m \\_______|\\_______/ │\n", - "│ \u001b[32m \u001b[0m $$\\ $$ |\u001b[31m\u001b[0m │\n", - "│ \u001b[32m \u001b[0m \\$$$$$$ |\u001b[31m\u001b[0m │\n", - "│ \u001b[32m \u001b[0m \\______/\u001b[31m\u001b[0m │\n", - "│ │\n", - "│ Inspect - Edit - Evolve Neural Networks │\n", - "│ │\n", - "╰─── By GrayBx, v1.3.3.dev8 ──────────────────────────────────────────────────────────────────────────╯\n", - "\n", - "\n", - "16/07/2026-09:29:29.690 INFO:weightslab.utils.telemetry:ping_import: WeightsLab uses anonymous usage data (package version used, and OS name). Set WL_NO_TELEMETRY=1 to disable.\n" - ] - } - ], + "outputs": [], "source": [ "# Install WeightsLab. Colab already ships torch, torchvision, numpy, scikit-learn\n", "# and Pillow, so nothing extra is needed here.\n", @@ -350,7 +166,6 @@ "import warnings\n", "warnings.filterwarnings(\"ignore\")\n", "\n", - "from weightslab.integrations.ultralytics import WLAwareSegmentationTrainer\n", "from ultralytics import YOLO\n", "\n", "# Hyperparameters registered with WeightsLab -> live-editable from Weights Studio.\n", @@ -375,7 +190,7 @@ "# Read back the (now live) hyperparameters and hand training to WLAwareSegmentationTrainer.\n", "model = YOLO(cfg[\"model\"])\n", "results = model.train(\n", - " trainer=WLAwareSegmentationTrainer,\n", + " trainer=wl.WLAwareSegmentationTrainer,\n", " data=str(cfg[\"data\"]),\n", " imgsz=cfg[\"image_size\"],\n", " epochs=int(cfg[\"epochs\"]),\n", @@ -405,7 +220,7 @@ }, { "cell_type": "code", - "execution_count": 3, + "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/", @@ -415,1230 +230,7 @@ "id": "uOJdXGA3Ks7o", "outputId": "26652fa3-6a15-43ce-9d75-034eaba21ccd" }, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "1241 samples tracked\n" - ] - }, - { - "data": { - "application/vnd.google.colaboratory.intrinsic+json": { - "repr_error": "Out of range float values are not JSON compliant: nan", - "type": "dataframe" - }, - "text/html": [ - "\n", - "
\n", - "
\n", - "\n", - "\n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - "
predictionprediction_rawtargetdiscardedorigintask_typelast_seengroup_idmember_rankimg_pathcls
sample_idannotation_id
00NoneNone[0.2335685, 0.254695, 0.4553995, 0.43075103, 1...Falsetrain_loader-1.000.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
10NoneNone[0.251174, 0.24061048, 0.44366202, 0.4307515, ...Falsetrain_loader-1.010.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
20NoneNoneNoneFalsetrain_loader-1.020.0/content/datasets/brain-tumor/images/train/000...[[1.0], [1.0]]
1NoneNone[0.523474, 0.3978875, 0.634976, 0.48122054, 1....NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.40610352, 0.415493, 0.5234745, 0.522301, 1....NoneNaNNoneNaNNoneNaNNoneNone
30NoneNone[0.403756, 0.3826295, 0.637324, 0.5152585, 1.0...Falsetrain_loader-1.030.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
40NoneNone[0.4436615, 0.3814555, 0.5938965, 0.4518785, 1...Falsetrain_loader-1.040.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
50NoneNone[0.6044605, 0.3990615, 0.67253554, 0.46596247,...Falsetrain_loader-1.050.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
60NoneNone[0.410798, 0.4436625, 0.556338, 0.51173747, 1....Falsetrain_loader-1.060.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
70NoneNone[0.650235, 0.4166665, 0.73826295, 0.48356748, ...Falsetrain_loader-1.070.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
80NoneNone[0.64788747, 0.388498, 0.7582165, 0.49413198, ...Falsetrain_loader-1.080.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
90NoneNone[0.6807515, 0.40140802, 0.75704247, 0.45774597...Falsetrain_loader-1.090.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
100NoneNone[0.431925, 0.1701875, 0.561033, 0.3180745, 1.0...Falsetrain_loader-1.0100.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
110NoneNone[0.42018852, 0.34507048, 0.5434275, 0.4941315,...Falsetrain_loader-1.0110.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
120NoneNone[0.3732395, 0.348592, 0.54812247, 0.46009403, ...Falsetrain_loader-1.0120.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
130NoneNone[0.37793398, 0.3814555, 0.48122, 0.45892054, 1...Falsetrain_loader-1.0130.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
140NoneNoneNoneFalsetrain_loader-1.0140.0/content/datasets/brain-tumor/images/train/000...[[1.0], [1.0]]
1NoneNone[0.43427247, 0.3556335, 0.51173747, 0.43075046...NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.4929575, 0.44718304, 0.5375585, 0.485915, 1...NoneNaNNoneNaNNoneNaNNoneNone
150NoneNoneNoneFalsetrain_loader-1.0150.0/content/datasets/brain-tumor/images/train/000...[[1.0], [1.0]]
1NoneNone[0.629108, 0.39788702, 0.67840403, 0.453051, 1...NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.4530515, 0.4107985, 0.58568054, 0.5105635, ...NoneNaNNoneNaNNoneNaNNoneNone
160NoneNoneNoneFalsetrain_loader-1.0160.0/content/datasets/brain-tumor/images/train/000...[[1.0], [1.0]]
1NoneNone[0.45070451, 0.3967135, 0.5762915, 0.5152585, ...NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.63732404, 0.403756, 0.684272, 0.46009403, 1...NoneNaNNoneNaNNoneNaNNoneNone
170NoneNone[0.43779355, 0.3826295, 0.61032856, 0.5223005,...Falsetrain_loader-1.0170.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
180NoneNone[0.4401405, 0.374413, 0.6068075, 0.529343, 1.0...Falsetrain_loader-1.0180.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
190NoneNone[0.536385, 0.4495305, 0.600939, 0.49999946, 0....Falsetrain_loader-1.0190.0/content/datasets/brain-tumor/images/train/000...[[0.0]]
200NoneNoneNoneFalsetrain_loader-1.0200.0/content/datasets/brain-tumor/images/train/000...[[0.0], [0.0], [0.0]]
1NoneNone[0.52347445, 0.4424885, 0.6150235, 0.5281695, ...NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.456573, 0.515258, 0.504695, 0.562206, 0.0, ...NoneNaNNoneNaNNoneNaNNoneNone
3NoneNone[0.504695, 0.44248852, 0.54342705, 0.48004752,...NoneNaNNoneNaNNoneNaNNoneNone
210NoneNone[0.54107946, 0.456573, 0.6103285, 0.511737, 0....Falsetrain_loader-1.0210.0/content/datasets/brain-tumor/images/train/000...[[0.0]]
220NoneNone[0.3298125, 0.2042255, 0.4225355, 0.2957745, 0...Falsetrain_loader-1.0220.0/content/datasets/brain-tumor/images/train/000...[[0.0]]
230NoneNone[0.517606, 0.18661949, 0.615024, 0.2875585, 1....Falsetrain_loader-1.0230.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
240NoneNone[0.51173747, 0.18779299, 0.6126765, 0.308685, ...Falsetrain_loader-1.0240.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
250NoneNoneNoneFalsetrain_loader-1.0250.0/content/datasets/brain-tumor/images/train/000...[[0.0], [0.0]]
1NoneNone[0.30985948, 0.3626765, 0.3861505, 0.43896753,...NoneNaNNoneNaNNoneNaNNoneNone
2NoneNone[0.41314548, 0.3427225, 0.48122054, 0.42018753...NoneNaNNoneNaNNoneNaNNoneNone
260NoneNone[0.25821602, 0.30046952, 0.46830997, 0.4577465...Falsetrain_loader-1.0260.0/content/datasets/brain-tumor/images/train/000...[[0.0]]
270NoneNone[0.2734745, 0.2828635, 0.44718352, 0.4577465, ...Falsetrain_loader-1.0270.0/content/datasets/brain-tumor/images/train/000...[[0.0]]
280NoneNone[0.5915495, 0.23591551, 0.6455405, 0.2887325, ...Falsetrain_loader-1.0280.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
290NoneNone[0.5903755, 0.231221, 0.6420185, 0.280517, 1.0...Falsetrain_loader-1.0290.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
300NoneNone[0.40493003, 0.2018775, 0.469484, 0.2699525, 1...Falsetrain_loader-1.0300.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
310NoneNone[0.39554, 0.19483551, 0.483568, 0.2887325, 1.0...Falsetrain_loader-1.0310.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
320NoneNone[0.4025825, 0.212441, 0.45539945, 0.269953, 1....Falsetrain_loader-1.0320.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
330NoneNone[0.2265255, 0.2136145, 0.3438965, 0.30751148, ...Falsetrain_loader-1.0330.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
340NoneNone[0.634976, 0.349765, 0.7723, 0.456573, 1.0, 1.0]Falsetrain_loader-1.0340.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
350NoneNone[0.6079815, 0.320423, 0.7969485, 0.489437, 1.0...Falsetrain_loader-1.0350.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
360NoneNone[0.6185445, 0.30164298, 0.7758215, 0.494131, 1...Falsetrain_loader-1.0360.0/content/datasets/brain-tumor/images/train/000...[[1.0]]
\n", - "
\n", - "
\n", - "\n", - "
\n", - " \n", - "\n", - " \n", - "\n", - " \n", - "
\n", - "\n", - "\n", - "
\n", - "
\n" - ], - "text/plain": [ - " prediction prediction_raw \\\n", - "sample_id annotation_id \n", - "0 0 None None \n", - "1 0 None None \n", - "2 0 None None \n", - " 1 None None \n", - " 2 None None \n", - "3 0 None None \n", - "4 0 None None \n", - "5 0 None None \n", - "6 0 None None \n", - "7 0 None None \n", - "8 0 None None \n", - "9 0 None None \n", - "10 0 None None \n", - "11 0 None None \n", - "12 0 None None \n", - "13 0 None None \n", - "14 0 None None \n", - " 1 None None \n", - " 2 None None \n", - "15 0 None None \n", - " 1 None None \n", - " 2 None None \n", - "16 0 None None \n", - " 1 None None \n", - " 2 None None \n", - "17 0 None None \n", - "18 0 None None \n", - "19 0 None None \n", - "20 0 None None \n", - " 1 None None \n", - " 2 None None \n", - " 3 None None \n", - "21 0 None None \n", - "22 0 None None \n", - "23 0 None None \n", - "24 0 None None \n", - "25 0 None None \n", - " 1 None None \n", - " 2 None None \n", - "26 0 None None \n", - "27 0 None None \n", - "28 0 None None \n", - "29 0 None None \n", - "30 0 None None \n", - "31 0 None None \n", - "32 0 None None \n", - "33 0 None None \n", - "34 0 None None \n", - "35 0 None None \n", - "36 0 None None \n", - "\n", - " target \\\n", - "sample_id annotation_id \n", - "0 0 [0.2335685, 0.254695, 0.4553995, 0.43075103, 1... \n", - "1 0 [0.251174, 0.24061048, 0.44366202, 0.4307515, ... \n", - "2 0 None \n", - " 1 [0.523474, 0.3978875, 0.634976, 0.48122054, 1.... \n", - " 2 [0.40610352, 0.415493, 0.5234745, 0.522301, 1.... \n", - "3 0 [0.403756, 0.3826295, 0.637324, 0.5152585, 1.0... \n", - "4 0 [0.4436615, 0.3814555, 0.5938965, 0.4518785, 1... \n", - "5 0 [0.6044605, 0.3990615, 0.67253554, 0.46596247,... \n", - "6 0 [0.410798, 0.4436625, 0.556338, 0.51173747, 1.... \n", - "7 0 [0.650235, 0.4166665, 0.73826295, 0.48356748, ... \n", - "8 0 [0.64788747, 0.388498, 0.7582165, 0.49413198, ... \n", - "9 0 [0.6807515, 0.40140802, 0.75704247, 0.45774597... \n", - "10 0 [0.431925, 0.1701875, 0.561033, 0.3180745, 1.0... \n", - "11 0 [0.42018852, 0.34507048, 0.5434275, 0.4941315,... \n", - "12 0 [0.3732395, 0.348592, 0.54812247, 0.46009403, ... \n", - "13 0 [0.37793398, 0.3814555, 0.48122, 0.45892054, 1... \n", - "14 0 None \n", - " 1 [0.43427247, 0.3556335, 0.51173747, 0.43075046... \n", - " 2 [0.4929575, 0.44718304, 0.5375585, 0.485915, 1... \n", - "15 0 None \n", - " 1 [0.629108, 0.39788702, 0.67840403, 0.453051, 1... \n", - " 2 [0.4530515, 0.4107985, 0.58568054, 0.5105635, ... \n", - "16 0 None \n", - " 1 [0.45070451, 0.3967135, 0.5762915, 0.5152585, ... \n", - " 2 [0.63732404, 0.403756, 0.684272, 0.46009403, 1... \n", - "17 0 [0.43779355, 0.3826295, 0.61032856, 0.5223005,... \n", - "18 0 [0.4401405, 0.374413, 0.6068075, 0.529343, 1.0... \n", - "19 0 [0.536385, 0.4495305, 0.600939, 0.49999946, 0.... \n", - "20 0 None \n", - " 1 [0.52347445, 0.4424885, 0.6150235, 0.5281695, ... \n", - " 2 [0.456573, 0.515258, 0.504695, 0.562206, 0.0, ... \n", - " 3 [0.504695, 0.44248852, 0.54342705, 0.48004752,... \n", - "21 0 [0.54107946, 0.456573, 0.6103285, 0.511737, 0.... \n", - "22 0 [0.3298125, 0.2042255, 0.4225355, 0.2957745, 0... \n", - "23 0 [0.517606, 0.18661949, 0.615024, 0.2875585, 1.... \n", - "24 0 [0.51173747, 0.18779299, 0.6126765, 0.308685, ... \n", - "25 0 None \n", - " 1 [0.30985948, 0.3626765, 0.3861505, 0.43896753,... \n", - " 2 [0.41314548, 0.3427225, 0.48122054, 0.42018753... \n", - "26 0 [0.25821602, 0.30046952, 0.46830997, 0.4577465... \n", - "27 0 [0.2734745, 0.2828635, 0.44718352, 0.4577465, ... \n", - "28 0 [0.5915495, 0.23591551, 0.6455405, 0.2887325, ... \n", - "29 0 [0.5903755, 0.231221, 0.6420185, 0.280517, 1.0... \n", - "30 0 [0.40493003, 0.2018775, 0.469484, 0.2699525, 1... \n", - "31 0 [0.39554, 0.19483551, 0.483568, 0.2887325, 1.0... \n", - "32 0 [0.4025825, 0.212441, 0.45539945, 0.269953, 1.... \n", - "33 0 [0.2265255, 0.2136145, 0.3438965, 0.30751148, ... \n", - "34 0 [0.634976, 0.349765, 0.7723, 0.456573, 1.0, 1.0] \n", - "35 0 [0.6079815, 0.320423, 0.7969485, 0.489437, 1.0... \n", - "36 0 [0.6185445, 0.30164298, 0.7758215, 0.494131, 1... \n", - "\n", - " discarded origin task_type last_seen group_id \\\n", - "sample_id annotation_id \n", - "0 0 False train_loader -1.0 0 \n", - "1 0 False train_loader -1.0 1 \n", - "2 0 False train_loader -1.0 2 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - "3 0 False train_loader -1.0 3 \n", - "4 0 False train_loader -1.0 4 \n", - "5 0 False train_loader -1.0 5 \n", - "6 0 False train_loader -1.0 6 \n", - "7 0 False train_loader -1.0 7 \n", - "8 0 False train_loader -1.0 8 \n", - "9 0 False train_loader -1.0 9 \n", - "10 0 False train_loader -1.0 10 \n", - "11 0 False train_loader -1.0 11 \n", - "12 0 False train_loader -1.0 12 \n", - "13 0 False train_loader -1.0 13 \n", - "14 0 False train_loader -1.0 14 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - "15 0 False train_loader -1.0 15 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - "16 0 False train_loader -1.0 16 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - "17 0 False train_loader -1.0 17 \n", - "18 0 False train_loader -1.0 18 \n", - "19 0 False train_loader -1.0 19 \n", - "20 0 False train_loader -1.0 20 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - " 3 None NaN None NaN None \n", - "21 0 False train_loader -1.0 21 \n", - "22 0 False train_loader -1.0 22 \n", - "23 0 False train_loader -1.0 23 \n", - "24 0 False train_loader -1.0 24 \n", - "25 0 False train_loader -1.0 25 \n", - " 1 None NaN None NaN None \n", - " 2 None NaN None NaN None \n", - "26 0 False train_loader -1.0 26 \n", - "27 0 False train_loader -1.0 27 \n", - "28 0 False train_loader -1.0 28 \n", - "29 0 False train_loader -1.0 29 \n", - "30 0 False train_loader -1.0 30 \n", - "31 0 False train_loader -1.0 31 \n", - "32 0 False train_loader -1.0 32 \n", - "33 0 False train_loader -1.0 33 \n", - "34 0 False train_loader -1.0 34 \n", - "35 0 False train_loader -1.0 35 \n", - "36 0 False train_loader -1.0 36 \n", - "\n", - " member_rank \\\n", - "sample_id annotation_id \n", - "0 0 0.0 \n", - "1 0 0.0 \n", - "2 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - "3 0 0.0 \n", - "4 0 0.0 \n", - "5 0 0.0 \n", - "6 0 0.0 \n", - "7 0 0.0 \n", - "8 0 0.0 \n", - "9 0 0.0 \n", - "10 0 0.0 \n", - "11 0 0.0 \n", - "12 0 0.0 \n", - "13 0 0.0 \n", - "14 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - "15 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - "16 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - "17 0 0.0 \n", - "18 0 0.0 \n", - "19 0 0.0 \n", - "20 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - " 3 NaN \n", - "21 0 0.0 \n", - "22 0 0.0 \n", - "23 0 0.0 \n", - "24 0 0.0 \n", - "25 0 0.0 \n", - " 1 NaN \n", - " 2 NaN \n", - "26 0 0.0 \n", - "27 0 0.0 \n", - "28 0 0.0 \n", - "29 0 0.0 \n", - "30 0 0.0 \n", - "31 0 0.0 \n", - "32 0 0.0 \n", - "33 0 0.0 \n", - "34 0 0.0 \n", - "35 0 0.0 \n", - "36 0 0.0 \n", - "\n", - " img_path \\\n", - "sample_id annotation_id \n", - "0 0 /content/datasets/brain-tumor/images/train/000... \n", - "1 0 /content/datasets/brain-tumor/images/train/000... \n", - "2 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - "3 0 /content/datasets/brain-tumor/images/train/000... \n", - "4 0 /content/datasets/brain-tumor/images/train/000... \n", - "5 0 /content/datasets/brain-tumor/images/train/000... \n", - "6 0 /content/datasets/brain-tumor/images/train/000... \n", - "7 0 /content/datasets/brain-tumor/images/train/000... \n", - "8 0 /content/datasets/brain-tumor/images/train/000... \n", - "9 0 /content/datasets/brain-tumor/images/train/000... \n", - "10 0 /content/datasets/brain-tumor/images/train/000... \n", - "11 0 /content/datasets/brain-tumor/images/train/000... \n", - "12 0 /content/datasets/brain-tumor/images/train/000... \n", - "13 0 /content/datasets/brain-tumor/images/train/000... \n", - "14 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - "15 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - "16 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - "17 0 /content/datasets/brain-tumor/images/train/000... \n", - "18 0 /content/datasets/brain-tumor/images/train/000... \n", - "19 0 /content/datasets/brain-tumor/images/train/000... \n", - "20 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - " 3 None \n", - "21 0 /content/datasets/brain-tumor/images/train/000... \n", - "22 0 /content/datasets/brain-tumor/images/train/000... \n", - "23 0 /content/datasets/brain-tumor/images/train/000... \n", - "24 0 /content/datasets/brain-tumor/images/train/000... \n", - "25 0 /content/datasets/brain-tumor/images/train/000... \n", - " 1 None \n", - " 2 None \n", - "26 0 /content/datasets/brain-tumor/images/train/000... \n", - "27 0 /content/datasets/brain-tumor/images/train/000... \n", - "28 0 /content/datasets/brain-tumor/images/train/000... \n", - "29 0 /content/datasets/brain-tumor/images/train/000... \n", - "30 0 /content/datasets/brain-tumor/images/train/000... \n", - "31 0 /content/datasets/brain-tumor/images/train/000... \n", - "32 0 /content/datasets/brain-tumor/images/train/000... \n", - "33 0 /content/datasets/brain-tumor/images/train/000... \n", - "34 0 /content/datasets/brain-tumor/images/train/000... \n", - "35 0 /content/datasets/brain-tumor/images/train/000... \n", - "36 0 /content/datasets/brain-tumor/images/train/000... \n", - "\n", - " cls \n", - "sample_id annotation_id \n", - "0 0 [[1.0]] \n", - "1 0 [[1.0]] \n", - "2 0 [[1.0], [1.0]] \n", - " 1 None \n", - " 2 None \n", - "3 0 [[1.0]] \n", - "4 0 [[1.0]] \n", - "5 0 [[1.0]] \n", - "6 0 [[1.0]] \n", - "7 0 [[1.0]] \n", - "8 0 [[1.0]] \n", - "9 0 [[1.0]] \n", - "10 0 [[1.0]] \n", - "11 0 [[1.0]] \n", - "12 0 [[1.0]] \n", - "13 0 [[1.0]] \n", - "14 0 [[1.0], [1.0]] \n", - " 1 None \n", - " 2 None \n", - "15 0 [[1.0], [1.0]] \n", - " 1 None \n", - " 2 None \n", - "16 0 [[1.0], [1.0]] \n", - " 1 None \n", - " 2 None \n", - "17 0 [[1.0]] \n", - "18 0 [[1.0]] \n", - "19 0 [[0.0]] \n", - "20 0 [[0.0], [0.0], [0.0]] \n", - " 1 None \n", - " 2 None \n", - " 3 None \n", - "21 0 [[0.0]] \n", - "22 0 [[0.0]] \n", - "23 0 [[1.0]] \n", - "24 0 [[1.0]] \n", - "25 0 [[0.0], [0.0]] \n", - " 1 None \n", - " 2 None \n", - "26 0 [[0.0]] \n", - "27 0 [[0.0]] \n", - "28 0 [[1.0]] \n", - "29 0 [[1.0]] \n", - "30 0 [[1.0]] \n", - "31 0 [[1.0]] \n", - "32 0 [[1.0]] \n", - "33 0 [[1.0]] \n", - "34 0 [[1.0]] \n", - "35 0 [[1.0]] \n", - "36 0 [[1.0]] " - ] - }, - "metadata": {}, - "output_type": "display_data" - }, - { - "data": { - "text/html": [ - "🔗 Open the full grid in Weights Studio" - ], - "text/plain": [ - "" - ] - }, - "metadata": {}, - "output_type": "display_data" - } - ], + "outputs": [], "source": [ "import pandas as pd\n", "from IPython.display import display, HTML\n", diff --git a/weightslab/examples/Notebooks/Usecases/wl-segmentation-loss-shapes-classification.ipynb b/weightslab/examples/Notebooks/Usecases/wl-segmentation-loss-shapes-classification.ipynb index 52c70781..13e1c39e 100644 --- a/weightslab/examples/Notebooks/Usecases/wl-segmentation-loss-shapes-classification.ipynb +++ b/weightslab/examples/Notebooks/Usecases/wl-segmentation-loss-shapes-classification.ipynb @@ -103,9 +103,6 @@ "from torch import optim\n", "from tqdm.auto import tqdm\n", "\n", - "from weightslab.components.global_monitoring import (\n", - " guard_training_context, guard_testing_context,\n", - ")\n", "from utils.data import BDD100kSegDataset, seg_collate\n", "from utils.model import SmallUNet\n", "from utils.criterions import (\n", @@ -336,7 +333,7 @@ "\n", "\n", "def train_step(loader, model, optimizer, sig):\n", - " with guard_training_context:\n", + " with wl.guard_training_context:\n", " inputs, ids, labels, _ = next(loader)\n", " inputs = inputs.to(device)\n", " labels = [[m.to(device) for m in insts] for insts in labels]\n", @@ -352,7 +349,7 @@ "\n", "def evaluate(loader, model, sig, n_batches):\n", " losses = dices = 0.0\n", - " with guard_testing_context, torch.no_grad():\n", + " with wl.guard_testing_context, torch.no_grad():\n", " for inputs, ids, labels, _ in loader:\n", " inputs = inputs.to(device)\n", " labels = [[m.to(device) for m in insts] for insts in labels]\n", diff --git a/weightslab/examples/Notebooks/Usecases/ws-segmentation-loss-shapes-classification.ipynb b/weightslab/examples/Notebooks/Usecases/ws-segmentation-loss-shapes-classification.ipynb index 14966922..25b0274c 100644 --- a/weightslab/examples/Notebooks/Usecases/ws-segmentation-loss-shapes-classification.ipynb +++ b/weightslab/examples/Notebooks/Usecases/ws-segmentation-loss-shapes-classification.ipynb @@ -89,9 +89,6 @@ "from tqdm.auto import tqdm\n", "\n", "import weightslab as wl\n", - "from weightslab.components.global_monitoring import (\n", - " guard_training_context, guard_testing_context,\n", - ")\n", "from utils.data import BDD100kSegDataset, seg_collate\n", "from utils.model import SmallUNet\n", "from utils.criterions import (\n", @@ -322,7 +319,7 @@ "\n", "\n", "def train_step(loader, model, optimizer, sig):\n", - " with guard_training_context:\n", + " with wl.guard_training_context:\n", " inputs, ids, labels, _ = next(loader)\n", " inputs = inputs.to(device)\n", " labels = [[m.to(device) for m in insts] for insts in labels]\n", @@ -338,7 +335,7 @@ "\n", "def evaluate(loader, model, sig, n_batches):\n", " losses = dices = 0.0\n", - " with guard_testing_context, torch.no_grad():\n", + " with wl.guard_testing_context, torch.no_grad():\n", " for inputs, ids, labels, _ in loader:\n", " inputs = inputs.to(device)\n", " labels = [[m.to(device) for m in insts] for insts in labels]\n", diff --git a/weightslab/examples/PyTorch/wl-ads-recommendation/main.py b/weightslab/examples/PyTorch/wl-ads-recommendation/main.py index 0c04b7c2..1ca8f9b3 100644 --- a/weightslab/examples/PyTorch/wl-ads-recommendation/main.py +++ b/weightslab/examples/PyTorch/wl-ads-recommendation/main.py @@ -29,10 +29,6 @@ from torchmetrics.classification import Accuracy import weightslab as wl -from weightslab.components.global_monitoring import ( - guard_training_context, - guard_testing_context, -) from utils.data import AdsCTRDataset, CATEGORICAL_CARDINALITIES, NUM_NUMERIC from utils.model import WideDeepCTR @@ -47,7 +43,7 @@ # ----------------------------------------------------------------------------- def train(loader, model, optimizer, criterion_mlt, device): """Single training step using the tracked dataloader + watched loss.""" - with guard_training_context: + with wl.guard_training_context: (inputs, ids, labels) = next(loader) inputs = inputs.to(device) labels = labels.to(device) @@ -75,7 +71,7 @@ def test(loader, model, criterion_mlt, metric_mlt, device, test_loader_len): losses = torch.tensor(0.0, device=device) for (inputs, ids, labels) in loader: - with guard_testing_context: + with wl.guard_testing_context: inputs = inputs.to(device) labels = labels.to(device) diff --git a/weightslab/examples/PyTorch/wl-ads-recommendation/verify_integration.py b/weightslab/examples/PyTorch/wl-ads-recommendation/verify_integration.py index b8b08d01..e2ff3c4d 100644 --- a/weightslab/examples/PyTorch/wl-ads-recommendation/verify_integration.py +++ b/weightslab/examples/PyTorch/wl-ads-recommendation/verify_integration.py @@ -32,10 +32,6 @@ sys.path.insert(0, os.path.dirname(__file__)) import weightslab as wl # noqa: E402 -from weightslab.components.global_monitoring import ( # noqa: E402 - guard_training_context, - guard_testing_context, -) from utils.data import ( # noqa: E402 AdsCTRDataset, @@ -100,7 +96,7 @@ def main() -> int: wl.start_training() for _ in range(120): - with guard_training_context: + with wl.guard_training_context: inputs, ids, labels = next(train_loader) optimizer.zero_grad() out = model(inputs) @@ -110,7 +106,7 @@ def main() -> int: optimizer.step() for inputs, ids, labels in test_loader: - with guard_testing_context: + with wl.guard_testing_context: out = model(inputs) preds = out.argmax(dim=1, keepdim=True) test_crit(out, labels, batch_ids=ids, preds=preds) diff --git a/weightslab/examples/PyTorch/wl-classification/main.py b/weightslab/examples/PyTorch/wl-classification/main.py index dc64f510..5741684e 100644 --- a/weightslab/examples/PyTorch/wl-classification/main.py +++ b/weightslab/examples/PyTorch/wl-classification/main.py @@ -25,10 +25,6 @@ import weightslab as wl from weightslab.examples.utils.baseline_models.pytorch.models import FashionCNN as CNN -from weightslab.components.global_monitoring import ( - guard_training_context, - guard_testing_context -) # Setup logging @@ -129,7 +125,7 @@ def __getitem__(self, idx): def train(loader, model, optimizer, criterion_mlt, device): """Single training step using the tracked dataloader + watched loss.""" - with guard_training_context: + with wl.guard_training_context: (inputs, ids, labels) = next(loader) inputs = inputs.to(device) labels = labels.to(device) @@ -165,7 +161,7 @@ def test(loader, model, criterion_mlt, metric_mlt, device, test_loader_len): losses = torch.tensor(0.0, device=device) for (inputs, ids, labels) in loader: - with guard_testing_context: + with wl.guard_testing_context: inputs = inputs.to(device) labels = labels.to(device) diff --git a/weightslab/examples/PyTorch/wl-clustering/main.py b/weightslab/examples/PyTorch/wl-clustering/main.py index a1f62461..40bb42fb 100644 --- a/weightslab/examples/PyTorch/wl-clustering/main.py +++ b/weightslab/examples/PyTorch/wl-clustering/main.py @@ -35,10 +35,6 @@ from face.data import FaceDataset from face.model import FaceEmbeddingModel from face.signals import FaceMetrics -from weightslab.components.global_monitoring import ( - guard_training_context, - guard_testing_context, -) logger = logging.getLogger(__name__) @@ -126,7 +122,7 @@ def train( while True: step += 1 - with guard_training_context: + with wl.guard_training_context: # ---- Fetch next batch (cycle loader) ---- try: images, batch_ids, labels, _metadata = next(data_iter) @@ -165,7 +161,7 @@ def train( ) if should_eval: - with guard_testing_context: + with wl.guard_testing_context: print(f"\n[eval@test] step {step}") metrics = evaluate(model=model, loader=test_loader, name="test") eval_history.append({"step": step, "metrics": metrics}) diff --git a/weightslab/examples/PyTorch/wl-detection/main.py b/weightslab/examples/PyTorch/wl-detection/main.py index d2122729..e283cdff 100644 --- a/weightslab/examples/PyTorch/wl-detection/main.py +++ b/weightslab/examples/PyTorch/wl-detection/main.py @@ -13,11 +13,6 @@ from torch import optim -from weightslab.components.global_monitoring import ( - guard_training_context, - guard_testing_context, -) - from utils.data import PennFudanDetectionDataset, det_collate from utils.model import SmallDetector from utils.criterions import ( @@ -43,7 +38,7 @@ def train(loader, model, optimizer, sig, device, grid_size, conf_thresh): DataSampleTrackingWrapper. `targets` is per sample a [N, 6] tensor of boxes ([x1, y1, x2, y2, class_id, confidence]); see utils/data.det_collate. """ - with guard_training_context: + with wl.guard_training_context: (inputs, ids, targets, _) = next(loader) inputs = inputs.to(device) targets = [t.to(device) for t in targets] @@ -73,7 +68,7 @@ def test(loader, model, sig, device, grid_size, conf_thresh, test_loader_len): """Full evaluation pass over the val loader.""" losses = 0.0 ious = 0.0 - with guard_testing_context, torch.no_grad(): + with wl.guard_testing_context, torch.no_grad(): for inputs, ids, targets, _ in loader: inputs = inputs.to(device) targets = [t.to(device) for t in targets] diff --git a/weightslab/examples/PyTorch/wl-fraud-detection/main.py b/weightslab/examples/PyTorch/wl-fraud-detection/main.py index 2213d7a3..1a02ccdb 100644 --- a/weightslab/examples/PyTorch/wl-fraud-detection/main.py +++ b/weightslab/examples/PyTorch/wl-fraud-detection/main.py @@ -37,10 +37,6 @@ from torchmetrics.classification import Precision, Recall, F1Score, AveragePrecision import weightslab as wl -from weightslab.components.global_monitoring import ( - guard_training_context, - guard_testing_context, -) from utils.data import load_creditcard_fraud, compute_class_weights, NUM_FEATURES from utils.model import FraudMLP @@ -60,7 +56,7 @@ def train(loader, model, optimizer, criterion_mlt, metrics, device): curves and per-sample signals as eval — logged under ``train-metric-*`` and ``train_metric/*`` — instead of only the loss. """ - with guard_training_context: + with wl.guard_training_context: (inputs, ids, labels) = next(loader) inputs = inputs.to(device) labels = labels.to(device) @@ -125,7 +121,7 @@ def test(loader, model, criterion_mlt, metrics, device, test_loader_len): losses = torch.tensor(0.0, device=device) for (inputs, ids, labels) in loader: - with guard_testing_context: + with wl.guard_testing_context: inputs = inputs.to(device) labels = labels.to(device) diff --git a/weightslab/examples/PyTorch/wl-fraud-detection/verify_integration.py b/weightslab/examples/PyTorch/wl-fraud-detection/verify_integration.py index 29acaf04..18fde638 100644 --- a/weightslab/examples/PyTorch/wl-fraud-detection/verify_integration.py +++ b/weightslab/examples/PyTorch/wl-fraud-detection/verify_integration.py @@ -37,10 +37,6 @@ sys.path.insert(0, os.path.dirname(__file__)) import weightslab as wl # noqa: E402 -from weightslab.components.global_monitoring import ( # noqa: E402 - guard_training_context, - guard_testing_context, -) from utils.data import ( # noqa: E402 FEATURE_NAMES, @@ -140,7 +136,7 @@ def main() -> int: # ---- drive a few real training steps ---- for _ in range(120): - with guard_training_context: + with wl.guard_training_context: inputs, ids, labels = next(train_loader) optimizer.zero_grad() out = model(inputs) @@ -151,7 +147,7 @@ def main() -> int: # ---- one full eval pass (populates test loss/prediction per sample) ---- for inputs, ids, labels in test_loader: - with guard_testing_context: + with wl.guard_testing_context: out = model(inputs) preds = out.argmax(dim=1, keepdim=True) probs = torch.softmax(out, dim=1)[:, 1] diff --git a/weightslab/examples/PyTorch/wl-image-generation/main.py b/weightslab/examples/PyTorch/wl-image-generation/main.py index e32f485d..fbe90b62 100644 --- a/weightslab/examples/PyTorch/wl-image-generation/main.py +++ b/weightslab/examples/PyTorch/wl-image-generation/main.py @@ -14,10 +14,6 @@ from torchmetrics.classification import BinaryAccuracy import weightslab as wl -from weightslab.components.global_monitoring import ( - guard_training_context, - guard_testing_context -) # Setup logging logging.basicConfig(level=logging.DEBUG) @@ -244,7 +240,7 @@ def flatten_lists(arr): def train_step(loader, model, optimizer, cls_criterion, contrastive_criterion, device, recon_weight, contrastive_weight): total_loss = None - with guard_training_context: + with wl.guard_training_context: try: images, ids, labels, metadata = next(loader) except StopIteration: @@ -327,7 +323,7 @@ def evaluate_all(loader, model, cls_criterion, contrastive_criterion, metric, de num_batches = 0 metric.reset() - with guard_testing_context, torch.no_grad(): + with wl.guard_testing_context, torch.no_grad(): for images, ids, labels, metadata in loader: inputs_flat = torch.cat([img.float() for img in images], dim=0).to(device) labels_flat = torch.cat([torch.tensor(l).float() if not isinstance(l, torch.Tensor) else l.float() for l in labels], dim=0).view(-1, 1).to(device) diff --git a/weightslab/examples/PyTorch/wl-segmentation/main.py b/weightslab/examples/PyTorch/wl-segmentation/main.py index fb84570f..7b5d86c7 100644 --- a/weightslab/examples/PyTorch/wl-segmentation/main.py +++ b/weightslab/examples/PyTorch/wl-segmentation/main.py @@ -13,11 +13,6 @@ from torch import optim -from weightslab.components.global_monitoring import ( - guard_training_context, - guard_testing_context, -) - from utils.data import BDD100kSegDataset, seg_collate from utils.model import SmallUNet from utils.criterions import ( @@ -80,7 +75,7 @@ def train(loader, model, optimizer, sig, device): loader yields (inputs, ids, labels, metadata) because of DataSampleTrackingWrapper. `labels` is per sample a LIST of instance masks (see utils/data.seg_collate). """ - with guard_training_context: + with wl.guard_training_context: (inputs, ids, labels, _) = next(loader) inputs = inputs.to(device) labels = [[m.to(device) for m in insts] for insts in labels] # per-sample list of instances @@ -111,7 +106,7 @@ def test(loader, model, sig, device, test_loader_len): """Full evaluation pass over the val loader.""" losses = 0.0 dices = 0.0 - with guard_testing_context, torch.no_grad(): + with wl.guard_testing_context, torch.no_grad(): for inputs, ids, labels, _ in loader: inputs = inputs.to(device) labels = [[m.to(device) for m in insts] for insts in labels] # per-sample list of instances diff --git a/weightslab/examples/PyTorch/wl-video-generation/main.py b/weightslab/examples/PyTorch/wl-video-generation/main.py index ffa7187b..b1725da2 100644 --- a/weightslab/examples/PyTorch/wl-video-generation/main.py +++ b/weightslab/examples/PyTorch/wl-video-generation/main.py @@ -39,10 +39,6 @@ import yaml import weightslab as wl -from weightslab.components.global_monitoring import ( - guard_training_context, - guard_testing_context, -) from utils.data import VideoGenerationDataset, AUDIO_SAMPLE_RATE from utils.model import FlowMatchingLoss, VideoFlowUNet, sample_clip @@ -191,7 +187,7 @@ def evaluate(loader, model, criterion, mode, device, cfg=None, attach=False, """ losses, count = 0.0, 0 for inputs, ids, labels, metadata in loader: - with guard_testing_context: + with wl.guard_testing_context: uids = list(metadata["uid"]) captions = list(metadata["caption"]) target, source = unpack_batch(inputs, mode, device) @@ -326,7 +322,7 @@ def evaluate(loader, model, criterion, mode, device, cfg=None, attach=False, for train_step in train_range: age = model.get_age() if hasattr(model, "get_age") else train_step - with guard_training_context: + with wl.guard_training_context: try: inputs, ids, labels, metadata = next(train_loader) except StopIteration: diff --git a/weightslab/examples/Ultralytics/wl-detection/main.py b/weightslab/examples/Ultralytics/wl-detection/main.py index 3174059f..53e5ab1a 100644 --- a/weightslab/examples/Ultralytics/wl-detection/main.py +++ b/weightslab/examples/Ultralytics/wl-detection/main.py @@ -23,7 +23,6 @@ ) import weightslab as wl -from weightslab.integrations.ultralytics import WLAwareTrainer from ultralytics import YOLO logging.getLogger("weightslab.watchdog.grpc_watchdog").setLevel(logging.ERROR) @@ -63,7 +62,7 @@ def main(): wl.start_training(timeout=3) # Blocks and keeps the main thread alive while background services run. Optionally set a timeout (seconds) to auto-stop. YOLO(model_name).train( - trainer=WLAwareTrainer, + trainer=wl.WLAwareTrainer, data=data_root, imgsz=image_size, epochs=1000 if max_steps == None else max(1, int(max_steps)), diff --git a/weightslab/examples/Usecases/wl-2d-lidar-detection/main.py b/weightslab/examples/Usecases/wl-2d-lidar-detection/main.py index 41c62365..ccd852a8 100644 --- a/weightslab/examples/Usecases/wl-2d-lidar-detection/main.py +++ b/weightslab/examples/Usecases/wl-2d-lidar-detection/main.py @@ -11,10 +11,6 @@ from torch import optim -from weightslab.components.global_monitoring import ( - guard_training_context, - guard_testing_context, -) from utils.data import Lidar2DDetectionDataset, lidar2d_collate, DEFAULT_PC_RANGE from utils.model import Pillars2DLite @@ -29,7 +25,7 @@ def train(loader, model, optimizer, sig, device, grid_size, pc_range, conf_thresh): - with guard_training_context: + with wl.guard_training_context: (points, ids, targets, _) = next(loader) points = points.to(device) targets = [t.to(device) for t in targets] @@ -47,7 +43,7 @@ def train(loader, model, optimizer, sig, device, grid_size, pc_range, conf_thres def test(loader, model, sig, device, grid_size, pc_range, conf_thresh, test_loader_len): losses = ious = 0.0 - with guard_testing_context, torch.no_grad(): + with wl.guard_testing_context, torch.no_grad(): for points, ids, targets, _ in loader: points = points.to(device) targets = [t.to(device) for t in targets] diff --git a/weightslab/examples/Usecases/wl-3d-lidar-detection/main.py b/weightslab/examples/Usecases/wl-3d-lidar-detection/main.py index ca18ad2e..1cac9552 100644 --- a/weightslab/examples/Usecases/wl-3d-lidar-detection/main.py +++ b/weightslab/examples/Usecases/wl-3d-lidar-detection/main.py @@ -12,10 +12,6 @@ from torch import optim -from weightslab.components.global_monitoring import ( - guard_training_context, - guard_testing_context, -) from utils.data import Lidar3DDetectionDataset, lidar_collate, DEFAULT_PC_RANGE from utils.model import PointPillarsLite @@ -99,7 +95,7 @@ def train(loader, model, optimizer, sig, device, grid_size, pc_range, conf_thres `targets` is per sample a [N, 9] tensor of 3D boxes ([cx, cy, cz, dx, dy, dz, yaw, class_id, confidence]); see utils/data. """ - with guard_training_context: + with wl.guard_training_context: (points, ids, targets, _) = next(loader) points = points.to(device) targets = [t.to(device) for t in targets] @@ -129,7 +125,7 @@ def test(loader, model, sig, device, grid_size, pc_range, conf_thresh, test_load """Full evaluation pass over the val loader.""" losses = 0.0 ious = 0.0 - with guard_testing_context, torch.no_grad(): + with wl.guard_testing_context, torch.no_grad(): for points, ids, targets, _ in loader: points = points.to(device) targets = [t.to(device) for t in targets] diff --git a/weightslab/examples/Usecases/wl-classification-signals_shape_classification/main.py b/weightslab/examples/Usecases/wl-classification-signals_shape_classification/main.py index b8bccc15..ebbaf31d 100644 --- a/weightslab/examples/Usecases/wl-classification-signals_shape_classification/main.py +++ b/weightslab/examples/Usecases/wl-classification-signals_shape_classification/main.py @@ -40,7 +40,6 @@ from collections import Counter -from weightslab import guard_training_context, guard_testing_context from utils.model import SmallCNN from utils.data import MNISTIdx @@ -134,7 +133,7 @@ def train(loader, model, opt, crit, dev, sync): for img, ids, lab in loader: img, lab = img.to(dev), lab.to(dev) sync(); ts = time.perf_counter() - with guard_training_context: + with wl.guard_training_context: opt.zero_grad() logits = model(img) crit(logits, lab, batch_ids=ids, preds=logits.argmax(1, keepdim=True)).mean().backward() @@ -151,7 +150,7 @@ def test(test_loader, model, crit, dev): with torch.no_grad(): for tb in test_loader: ti, tid, tl = tb[0].to(dev), tb[1], tb[2].to(dev) - with guard_testing_context: + with wl.guard_testing_context: tlg = model(ti) crit(tlg, tl, batch_ids=tid, preds=tlg.argmax(1, keepdim=True)) diff --git a/weightslab/examples/Usecases/wl-fashion-mnist-signals/main.py b/weightslab/examples/Usecases/wl-fashion-mnist-signals/main.py index ccaf838d..6f45e46d 100644 --- a/weightslab/examples/Usecases/wl-fashion-mnist-signals/main.py +++ b/weightslab/examples/Usecases/wl-fashion-mnist-signals/main.py @@ -71,10 +71,6 @@ from torchvision import datasets, transforms import weightslab as wl -from weightslab.components.global_monitoring import ( - guard_testing_context, - guard_training_context, -) logging.basicConfig(level=logging.ERROR) logger = logging.getLogger(__name__) @@ -206,7 +202,7 @@ def print_layer_legend(model): # ----------------------------------------------------------------------------- def train(loader, model, optimizer, criterion, device): """One training step. Nothing here logs model signals -- the hooks do.""" - with guard_training_context: + with wl.guard_training_context: inputs, ids, labels = next(loader) inputs = inputs.to(device) labels = labels.to(device) @@ -235,7 +231,7 @@ def test(loader, model, criterion, metric, device, num_batches): losses = torch.tensor(0.0, device=device) for inputs, ids, labels in loader: - with guard_testing_context, torch.no_grad(): + with wl.guard_testing_context, torch.no_grad(): inputs = inputs.to(device) labels = labels.to(device) diff --git a/weightslab/examples/Usecases/ws-signals-mnist/main.py b/weightslab/examples/Usecases/ws-signals-mnist/main.py index 867e188f..f9313a3a 100644 --- a/weightslab/examples/Usecases/ws-signals-mnist/main.py +++ b/weightslab/examples/Usecases/ws-signals-mnist/main.py @@ -24,7 +24,6 @@ from torchvision import datasets, transforms import weightslab as wl -from weightslab.components.global_monitoring import guard_training_context, guard_testing_context LOSS = "loss_sample" OUT = os.environ.get("WL_STRESS_OUT", "/tmp/wl_stress") @@ -134,7 +133,7 @@ def test_eval(): with torch.no_grad(): for tb in test_loader: ti, tid, tl = tb[0].to(dev), tb[1], tb[2].to(dev) - with guard_testing_context: + with wl.guard_testing_context: tlg = model(ti) crit(tlg, tl, batch_ids=tid, preds=tlg.argmax(1, keepdim=True)) @@ -148,7 +147,7 @@ def test_eval(): for img, ids, lab in loader: img, lab = img.to(dev), lab.to(dev) sync(); ts = time.perf_counter() - with guard_training_context: + with wl.guard_training_context: opt.zero_grad() logits = model(img) # only per-step call: the watched loss logs loss_sample and fires diff --git a/weightslab/integrations/ultralytics/README.md b/weightslab/integrations/ultralytics/README.md index ef9e464c..0bdbad28 100644 --- a/weightslab/integrations/ultralytics/README.md +++ b/weightslab/integrations/ultralytics/README.md @@ -14,13 +14,12 @@ changes to the model or to UL's training loop. ```python import weightslab as wl from ultralytics import YOLO -from weightslab.integrations.ultralytics import WLAwareTrainer # or WLAwareSegmentationTrainer wl.watch_or_edit(cfg, flag="hyperparameters", defaults=cfg) wl.serve() -YOLO("yolo11n.pt").train( # yolo11n-seg.pt for segmentation - trainer=WLAwareTrainer, # WLAwareSegmentationTrainer for segmentation +YOLO("yolo11n.pt").train( # yolo11n-seg.pt for segmentation + trainer=wl.WLAwareTrainer, # wl.WLAwareSegmentationTrainer for segmentation data="my_dataset.yaml", imgsz=640, epochs=100, batch=16, project="./logs", name="exp", workers=0, amp=False, diff --git a/weightslab/integrations/ultralytics/__init__.py b/weightslab/integrations/ultralytics/__init__.py index 3ec79271..36ce997b 100644 --- a/weightslab/integrations/ultralytics/__init__.py +++ b/weightslab/integrations/ultralytics/__init__.py @@ -16,13 +16,12 @@ import weightslab as wl from ultralytics import YOLO - from weightslab.integrations.ultralytics import WLAwareTrainer wl.watch_or_edit(cfg, flag="hyperparameters", defaults=cfg) wl.serve() YOLO(cfg["model"]).train( - trainer=WLAwareTrainer, # or WLAwareSegmentationTrainer + trainer=wl.WLAwareTrainer, # or wl.WLAwareSegmentationTrainer data=cfg["data_root"], imgsz=640, epochs=1000, batch=4, project="./logs", name="exp", # → WL log_dir/name workers=0, # WL invariant (parent-process uid counter) From c25bd45805ebde7afa5deca253781fbb2e6e9fd4 Mon Sep 17 00:00:00 2001 From: Guillaume Date: Thu, 10 Sep 2026 14:21:15 +0200 Subject: [PATCH 08/29] fix(notebook): keep the trainer's stdout out of notebook cell outputs (#305) The studio notebook runs an embedded ipykernel inside the trainer's own process, and IPKernelApp.initialize() swaps sys.stdout/sys.stderr process-wide for an OutStream that publishes every write to iopub. So the training loop's tqdm bar -- a different thread, writing continuously -- surfaced in whatever cell was last executed ("Training: 193497 steps ... train_loss=1.4612" in a cell that never asked for it). The legacy in-process kernel leaked the same way: contextlib.redirect_stdout is process-global too. Route the streams per write instead of per process: * _ThreadRoutedStream wraps ipykernel's OutStreams -- the thread currently running a cell reaches the kernel stream, every other thread (at any time, including while the kernel is idle) gets the console stream the process had before the kernel existed. Ownership comes from the pre_execute/post_execute hooks, which run on the real execution thread, so nothing here assumes which thread ipykernel picked for the shell channel. * capture_fd_output=False, or ipykernel's fd 1/2 pipe would re-capture exactly the writes that were just routed back to the terminal. The cost is that output written straight to the fds by C extensions no longer reaches the notebook -- for an in-process kernel sharing a terminal with the trainer, that is the better trade. * _LiveStream (legacy kernel) got the same thread check, falling back to the pre-redirect stream so those writes still reach the console. Two tests: one in the shared contract class, so both kernels are held to it (a background thread's output must not appear in a cell, while the cell's own print still does), and a legacy-only one proving the other thread's writes reach the console rather than being dropped. Known consequence: output from a thread a cell itself spawns now goes to the terminal, not the cell. Co-authored-by: Claude Opus 5 (1M context) --- .../services/test_notebook_service_unit.py | 62 ++++++++ .../trainer/services/notebook_service.py | 149 +++++++++++++++++- 2 files changed, 204 insertions(+), 7 deletions(-) diff --git a/tests/trainer/services/test_notebook_service_unit.py b/tests/trainer/services/test_notebook_service_unit.py index b02974f6..b8a6adfa 100644 --- a/tests/trainer/services/test_notebook_service_unit.py +++ b/tests/trainer/services/test_notebook_service_unit.py @@ -5,6 +5,8 @@ style of test_agent_service_unit.py. """ +import io +import sys import json import time import tempfile @@ -100,6 +102,37 @@ def test_stdout_streams_live_not_buffered_until_cell_finishes(self): f"the end ({total:.3f}s) -- looks buffered, not streamed live", ) + def test_other_threads_output_does_not_land_in_the_cell(self): + # Regression: the kernel shares its process with the trainer, and + # both kernels swap sys.stdout process-wide (redirect_stdout() in the + # legacy one, ipykernel's OutStream in the embedded one) -- so a + # training loop's tqdm bar, its own thread writing the whole time, + # surfaced in the output of whatever cell happened to be running. + stop = threading.Event() + + def _trainer(): + # Resolved per write, exactly as print() does, so this thread sees + # the kernel's swapped stream while a cell is running. + while not stop.is_set(): + print("Training: 193497 steps | train_loss=1.4612", file=sys.stdout) + time.sleep(0.01) + + noise = threading.Thread(target=_trainer, daemon=True) + noise.start() + try: + chunks = _run( + self.service, + "import time\nprint('cell-own-output')\ntime.sleep(0.3)", + ) + finally: + stop.set() + noise.join(timeout=2) + + outs = "".join(c.stdout for c in chunks if c.WhichOneof("payload") == "stdout") + self.assertIn("cell-own-output", outs, "the cell's own print was lost") + self.assertNotIn("Training:", outs, + "another thread's output leaked into the cell") + def test_interrupt_reports_false_when_nothing_running(self): resp = self.service.InterruptNotebookCell(pb2.InterruptNotebookCellRequest(), None) self.assertFalse(resp.ok) @@ -310,6 +343,35 @@ class TestNotebookKernelLegacy(_NotebookKernelContractTests, unittest.TestCase): def _make_service(self): return NotebookService(_fake_data_service(), root_log_dir=str(self.root)) + def test_other_threads_output_still_reaches_the_console(self): + # The flip side of the contract test above: writes the cell may not + # publish are handed to the stream that was in place before + # redirect_stdout(), not dropped on the floor -- the trainer's logs + # must keep showing up in the terminal while a cell runs. + console = io.StringIO() + stop = threading.Event() + + def _trainer(): + while not stop.is_set(): + print("Training: 193497 steps", file=sys.stdout) + time.sleep(0.01) + + real_stdout = sys.stdout + sys.stdout = console + noise = threading.Thread(target=_trainer, daemon=True) + noise.start() + try: + chunks = _run(self.service, "import time\ntime.sleep(0.3)") + finally: + stop.set() + noise.join(timeout=2) + sys.stdout = real_stdout + + outs = "".join(c.stdout for c in chunks if c.WhichOneof("payload") == "stdout") + self.assertNotIn("Training:", outs) + self.assertIn("Training:", console.getvalue(), + "the other thread's output never reached the console") + @unittest.skipUnless(_IPYKERNEL_AVAILABLE, "ipykernel/jupyter_client not installed") class TestNotebookKernelEmbedded(_NotebookKernelContractTests, unittest.TestCase): diff --git a/weightslab/trainer/services/notebook_service.py b/weightslab/trainer/services/notebook_service.py index de0d2e54..5679d741 100644 --- a/weightslab/trainer/services/notebook_service.py +++ b/weightslab/trainer/services/notebook_service.py @@ -28,6 +28,7 @@ import os import re import ast +import sys import json import time import ctypes @@ -409,6 +410,87 @@ def get_embedded_kernel_connection_file(wait_timeout: float = 8.0): return _EMBED_STATE["connection_file"] +# Ident of the thread currently running a notebook cell, or None while the +# kernel is idle. Set/cleared by the pre_execute/post_execute hooks, which run +# on the kernel's own execution thread -- so this needs no assumption about +# which thread ipykernel picked for the shell channel. +_CELL_THREAD = {"ident": None} + + +class _ThreadRoutedStream: + """stdout/stderr proxy that lets only the *cell's own* thread reach the + notebook; every other thread keeps writing to the real console. + + The embedded kernel shares its process with the trainer, and ipykernel's + ``init_io()`` swaps ``sys.stdout``/``sys.stderr`` process-wide for an + OutStream that ships everything to iopub. So a training loop's tqdm bar -- + another thread entirely, writing continuously -- landed in whichever cell + was last executed ("Training: 193497 steps ... train_loss=..." showing up + in a cell that never asked for it). Routing happens per write instead: + while a cell runs, its own thread reaches the kernel stream; anything else, + at any time, goes to the stream the process would have had without a + kernel. + """ + + def __init__(self, kernel_stream, console_stream): + self._kernel = kernel_stream + self._console = console_stream + + def _target(self): + ident = _CELL_THREAD["ident"] + if ident is not None and ident == threading.get_ident(): + return self._kernel + return self._console if self._console is not None else self._kernel + + def write(self, s): + return self._target().write(s) + + def writelines(self, lines): + target = self._target() + for line in lines: + target.write(line) + + def flush(self): + for stream in (self._kernel, self._console): + if stream is None: + continue + try: + stream.flush() + except Exception: # noqa: BLE001 -- a closed console must not break a cell + pass + + # Routed too, not delegated: tqdm asks isatty() once, when it is built, and + # a bar built on the training thread must get the console's answer ( + # refreshes) rather than the kernel OutStream's flat False. + def isatty(self): + try: + return bool(self._target().isatty()) + except Exception: # noqa: BLE001 + return False + + def fileno(self): + return self._target().fileno() + + def writable(self): + return True + + def __getattr__(self, name): + # encoding, errors, buffer, _original_stdstream_copy, ... -- whatever + # ipykernel or a library reaches for beyond the file protocol above. + return getattr(self._kernel, name) + + +def _install_thread_routed_streams(console_stdout, console_stderr) -> None: + """Wrap ipykernel's OutStreams so only cell threads publish to the notebook. + + Call after ``IPKernelApp.initialize()`` (which installs the OutStreams) and + before ``app.start()``. + """ + import sys as _sys + _sys.stdout = _ThreadRoutedStream(_sys.stdout, console_stdout) + _sys.stderr = _ThreadRoutedStream(_sys.stderr, console_stderr) + + def _run_embedded_kernel(connection_file: Path) -> None: import asyncio from ipykernel.kernelapp import IPKernelApp @@ -424,7 +506,19 @@ def _run_embedded_kernel(connection_file: Path) -> None: ns = build_notebook_namespace( _ACTIVE_BINDING["data_service"], _ACTIVE_BINDING["root_log_dir"]) + # The real console streams, grabbed before initialize() swaps them for + # ipykernel's OutStream -- _ThreadRoutedStream hands every non-cell thread + # back to these. + console_stdout, console_stderr = sys.stdout, sys.stderr + app = IPKernelApp.instance(connection_file=str(connection_file), matplotlib="inline") + # Without this, ipykernel replaces fd 1/2 with a pipe it forwards to iopub, + # which would swallow the console writes _ThreadRoutedStream routes back to + # the terminal (and re-publish the trainer's output into a cell anyway). + # The cost is that output written straight to the fds by C extensions no + # longer reaches the notebook -- for an in-process kernel sharing a + # terminal with the trainer, that is the better trade. + app.capture_fd_output = False # IPKernelApp.initialize() installs a SIGINT handler, and signal handlers # can only be installed on the main thread -- which an embedded kernel is # never on. ipykernel catches the resulting ValueError but logs it as @@ -482,6 +576,8 @@ def _run_embedded_kernel(connection_file: Path) -> None: if hasattr(_stream, "flush_interval"): _stream.flush_interval = 0.05 _install_kernel_hooks(app.shell) + # After initialize() (OutStreams exist), before start() (cells run). + _install_thread_routed_streams(console_stdout, console_stderr) logger.info("Embedded Jupyter kernel connection file: %s", connection_file) app.start() # blocks this thread forever (event loop) except Exception: @@ -496,6 +592,9 @@ def _install_kernel_hooks(shell) -> None: box = {"guard_cm": None} def _pre_execute(): + # This hook runs on the thread that executes the cell -- the one + # _ThreadRoutedStream lets through to the notebook. + _CELL_THREAD["ident"] = threading.get_ident() try: shell.user_ns["df"] = get_df(_ACTIVE_BINDING["data_service"]) except Exception: @@ -505,6 +604,7 @@ def _pre_execute(): box["guard_cm"] = cm def _post_execute(): + _CELL_THREAD["ident"] = None cm = box.pop("guard_cm", None) if cm is not None: cm.__exit__(None, None, None) @@ -639,19 +739,46 @@ class _LiveStream: """Write-only file-like object that forwards each write directly to ``emit(kind, text)`` instead of buffering into a StringIO -- lets stdout/ stderr reach the gRPC client as the cell actually prints, rather than only - after the whole cell finishes.""" + after the whole cell finishes. + + Only writes from the kernel worker thread are forwarded. redirect_stdout() + swaps ``sys.stdout`` for the whole process, and this kernel shares its + process with the trainer -- so without the thread check a training loop's + tqdm bar ends up in the output of whatever cell happens to be running. + Other threads keep writing to ``console``, the stream that was in place + before the redirect. + """ - def __init__(self, kind: str, emit): + def __init__(self, kind: str, emit, console=None, owner=None): self._kind = kind self._emit = emit + self._console = console + self._owner = owner if owner is not None else threading.get_ident() def write(self, s): - if s: - self._emit(self._kind, _capped(self._kind, s)) + if not s: + return 0 + if threading.get_ident() != self._owner: + if self._console is not None: + return self._console.write(s) + return len(s) + self._emit(self._kind, _capped(self._kind, s)) return len(s) def flush(self): - pass + if self._console is not None: + try: + self._console.flush() + except Exception: # noqa: BLE001 -- a closed console must not break a cell + pass + + def isatty(self): + if threading.get_ident() != self._owner and self._console is not None: + try: + return bool(self._console.isatty()) + except Exception: # noqa: BLE001 + return False + return False # --------------------------------------------------------------------------- @@ -771,10 +898,18 @@ def _run_on_kernel_thread(self, code: str, emit): except Exception: pass + # Captured before the redirect so _LiveStream can hand other + # threads' writes (the trainer's, typically) back to the console + # instead of publishing them into this cell's output. + console_stdout, console_stderr = sys.stdout, sys.stderr + owner = threading.get_ident() + try: with _WriteGuard.enforce(self._root_log_dir): - with contextlib.redirect_stdout(_LiveStream("stdout", emit)), \ - contextlib.redirect_stderr(_LiveStream("stderr", emit)): + with contextlib.redirect_stdout( + _LiveStream("stdout", emit, console_stdout, owner)), \ + contextlib.redirect_stderr( + _LiveStream("stderr", emit, console_stderr, owner)): result_repr = self._exec_with_last_expr(code) except BaseException: # noqa: BLE001 -- surface any user error (incl. an # interrupt() -injected KeyboardInterrupt) as a cell error, not a crash. From b0347d315963948d3c30c4876b961a06b84dc379 Mon Sep 17 00:00:00 2001 From: AlexGrayBox Date: Thu, 10 Sep 2026 14:28:23 +0200 Subject: [PATCH 09/29] Interactivity on 100GB+ datasets: O(change) view sync and ledger writes (#303) * docs(perf): register O(data) operations blocking 100GB+ interactivity Storage and serving paths whose cost scales with dataset size rather than with what changed. Storage findings re-verified against dev; two serving costs noted as already fixed upstream so they are not re-claimed as wins. No code changes - baseline and measurement protocol only. * docs(perf): triage every _slowUpdateInternals call site 16 of 18 sites only need fresh values for dirty rows (O(change)); only first build and schema change need a full reconstruction. Records why ApplyDataQuery filter paths stay on the rebuild (_is_filtered semantics) and why a no-client benchmark cannot show the difference. * perf(interactivity): make view refresh and ledger writes O(change) Every path that kept the served view in sync ran O(dataset): a full materialized-view rebuild on each signal tick, and a read-modify-append of the whole H5 table per upsert. At 3.96M rows that meant a 670s startup and lock holds long enough that the UI looked hung while training. Three changes, each turning a whole-dataset pass into a per-change one: data_service: differential view refresh (_fastUpdateInternals). The materialized view only ever recomputes index-derived state, so a value-only delta can be written straight into the existing view through a sample_id -> row-position map instead of rebuilding it. Falls back to the full path on any structural change (unknown sample_id, missing pos map, backlog over max_dirty), so correctness never depends on the fast path being right about a schema change. 9 call sites routed here; the 10 that genuinely change shape still force a rebuild. Off with WL_FAST_VIEW=0. The sort path is restructured to do its work off-lock: ops and the pos-map rebuild both run on a shallow copy, and only the pointer swap happens under the lock. Skipping the pos-map rebuild after a sort would have been a data corruption bug -- sorting reorders the view, so stale positions send differential writes to the wrong rows. h5_dataframe_store: in-place row updates via modify_coordinates, with a cached sample_id -> coordinate map (stable row positions come free from the no-row-loss invariant). PyTables cannot invalidate a column index during modify_coordinates, so _try_inplace refuses indexed tables and the caller falls back to the append path. Index construction is also split from storage layout: data_columns=True keeps the on-disk layout queryable while index=False keeps the flush path from rebuilding an index no hot-path read uses (92.7s -> 6.9s per upsert at 4M rows). dataframe_manager: replaces the per-row iterrows() scans that dominated startup with column-wise vectorised passes, and adds the dirty-row/source-row accessors the differential refresh needs. Also fixes an unrelated thumbnail bug in trainer_tools.process_sample: it unpacked exactly 3 values from _getitem_raw, whose contract is (data, id, target, *metadata). Any dataset implementing get_items() with metadata raised "too many values to unpack" and every cell in the grid came back with no image. Now unpacked positionally. Measured on 3.96M-row UltraEdit, A10G: startup 670s -> 250s H5 upsert (24 rows) 127.9s -> ~16ms snapshot flush 338s -> 0 samples max lock hold 129,440ms -> none over 1s throughput under UI 13% -> 45-51% of idle The residual loss under load is CPU/GIL contention (8 vCPUs shared by 6 dataloader workers, training, and image encode), not lock waiting. Known gaps, deliberately left for review: - ensure_index() has no caller yet, and conflicts with _try_inplace, which refuses indexed tables. It documents the deliberate-index path but is dead code as committed. - _POSMAP_CACHE is class-level and never evicts (~400-600MB at 4M rows). - Three ApplyDataQuery sites still force a full rebuild pending a decision on _is_filtered semantics. Co-Authored-By: Claude Opus 5 * fix(view): stop the served view silently diverging from the ledger The ledger was always correct; the view readers see was not, and every failure mode reported itself as success. On a 3.96M-sample run the UI showed ~2k samples with loss data at 34k steps, and toggling image modalities showed the source image twice. View correctness: * Address differential-sync rows by the SAMPLE_ID index level. The view is indexed (origin, sample_id), so get_level_values(0) returned origin and its intersection with the dirty sample_ids was always empty -- the sync wrote nothing and returned True, which suppressed the rebuild that would have repaired it. Positions now come from Index.get_indexer (cached hash engine), so it stays vectorised. * Rebuild when the ledger gains columns the view lacks. Per-sample signal columns are created on their first write, so on a fresh ledger the view predates them and could never gain them: sorting and histogramming failed with "column not in view" and last_seen served -1 forever. Checked before the dirty-set drain, since a schema gain is independent of dirty rows. * Keep the view-dirty backlog until a rebuild actually lands. It was discarded on overflow assuming the caller would rebuild, but the force path returns early on a contended lock -- those ids were then lost with nothing left to re-mark them. Cleared at the atomic view swap instead. * Log view-build failures as errors. They were swallowed at debug level and returned the previous view, so a broken build was indistinguishable from "no new data". Named image views: * Probe extra_images() on the unwrapped dataset -- WL's tracking wrapper does not forward it, so every named view was silently dropped. * Honour stats_to_retrieve for image views, but still advertise filtered-out views with an empty thumbnail so their toggles do not vanish from the panel. Cost: * Loss-shape autotagging runs on its own interval (WL_LOSS_SHAPE_INTERVAL_SECONDS, default 60s) instead of the 2s flush tick, where each pass cost ~990ms of GIL-held pandas work. * Signal-DAG history reads a bounded in-memory tail (WL_HISTORY_TAIL, default 16) instead of scanning per_sample -- 140ms per step at 20M rows, growing without bound. Neither an index nor a rewritten IN clause helped (1.1x/1.4x). * Skip array normalisation for columns the H5 write excludes: with predictions off it rasterised via get_mask, which reads the source image, for data that is never persisted. * Close inherited HDF5 fds in forked dataloader workers; they made HDF5 refuse the parent's read-write open and killed ledger persistence for 12 hours. Measured on the UltraEdit harness (859M params, batch 24, A10G) against an identical run with weightslab stubbed out: 1574ms -> ~1290ms/step versus a 1171ms baseline, i.e. +34% -> ~+10%. optrace.py is included: the @traced/hit markers the other files import are what located the sample_id level bug. Co-Authored-By: Claude Opus 5 * fix(histogram): bin over rows that carry a value, not every row in the view The numeric path cut bin boundaries by row position across the WHOLE view and dropped non-finite values only afterwards, so every bucket spanned len(view)/max_bins rows regardless of the data. On a 3,963,189-row view with 512 bins that is 7,740 rows per bucket -- so a column where only ~33k samples carry a value (any signal early in a run) collapsed into the first four buckets, and the remaining 500 sliced empty space into 1-6 sample slivers. Visible as four fat bars followed by a long tail of random-looking spikes, with the same 7740/7741 counts appearing on unrelated columns because the number came from the row count, not the data. These bars are a search surface over the loss landscape: each should be a click-target holding a comparable number of samples. Mask first, then cut equal-population boundaries over the finite subset. Row ORDER is untouched, so this is still "bin the current view by row order" -- it just stops counting rows that have nothing to show. Also fixes the per-(origin, discarded) sub-bars, which were grouped by the same positional bins. Co-Authored-By: Claude Opus 5 * fix(ledger): make the NB_SEEN lookup O(batch), and regenerate the protos Three things, all needed to make the branch runnable on top of dev at 4M rows. 1. Regenerated experiment_service_pb2{,_grpc}.py dev's .proto declares AnnotationExportFormat / EXPORT_FORMAT_CVAT but the committed gencode predates it, so a clean checkout of dev does not import at all: AttributeError: module 'weightslab.proto.experiment_service_pb2' has no attribute 'EXPORT_FORMAT_CVAT' Regenerated with grpcio-tools 1.68.1 (protoc 5.28.1), matching the runtime version already pinned in the file, so the gencode major does not move. 2. Cache the level-0 index for sample-id coercion _coerce_sample_id_for_index() called index.get_level_values(0) on every invocation. That materialises a fresh Index over all rows, and a fresh Index carries a fresh hash engine, so each `sid in level_0_values` paid a full engine build -- twice per sample when the int probe missed. enqueue_batch does that per sample, 24x a step: training sat at 0 iterations with the main thread pinned at 100% CPU inside pandas __contains__ (py-spy: active+gil). Cached on the index object's identity. pandas Index is immutable, so any reindex or rebuild yields a new object and invalidates it; membership semantics are unchanged, the engine is simply reused. 3. Positional NB_SEEN lookup get_sample_column_values() then still materialised both index levels, ran isin over every row and copied a boolean-masked frame -- a full pass plus a copy over 3.96M rows to read 24 integers, ~1.2s/step (signals 6ms -> 1230ms, total 1290ms -> 2500ms). The wanted rows are exactly (sample_id, 0), so resolve their positions with Index.get_indexer instead. Falls back to the original scan when the index is not unique. Measured on the UltraEdit harness (859M params, batch 24, A10G, 3.96M samples): signals 1230ms -> 28-41ms total 2500ms -> 1310-1338ms (1171ms with weightslab stubbed out) NB_SEEN now actually increments (it was stuck at 0 before dev's fix), verified against the ledger: rows with nb_seen>0 equals rows with last_seen>=0. The UI contract suite passes 21/21 on this build. Co-Authored-By: Claude Opus 5 * fix(data_service): bin the numeric histogram over the whole view again Each bar must cover total_rows / max_bins samples so the chart carries density: a column that is only 0.2% populated should show a few filled bars and the rest empty. Binning over just the rows that carry a value made the chart look equally full at any coverage, which reads as "every sample already has a loss". Co-Authored-By: Claude Opus 5 * proto: regenerate with package-relative imports after the dev merge dev checks in generated code that does a flat 'import experiment_service_pb2', which only resolves if weightslab/proto is itself on sys.path. Imported as a package -- which is how the trainer loads it -- startup dies with ModuleNotFoundError. Regenerated from the merged .proto at the repo root so dev's new RPCs are kept and the import is package-relative again. Co-Authored-By: Claude Opus 5 * fix(logger): import deque alongside defaultdict The merge re-applied our in-memory history tail onto dev logger.py, which imports only defaultdict. Every per-sample write then raised NameError inside _stage_sample_row. The caller swallows per-signal exceptions, so nothing crashed: the tail just stayed empty, sig/loss_debiased failed every step, and loss_shape had no history to classify. Co-Authored-By: Claude Opus 5 * fix(shapes): keep the label cache on top of dev write_signal_shapes dev rewrite keeps the O(change) read and adds exp_hash scoping, both kept. What it dropped is the label cache, which two behaviours depended on: - an incremental pass still returns a distribution over the WHOLE dataset, not just the samples it happened to touch; - a sample whose label did not change is not re-written to the ledger. Both are asserted by e2e_autotag (distribution_covers_dataset, incremental_writes_bounded), which failed on the merge until this went back. The test also moves to dev parameter name, sample_ids. Co-Authored-By: Claude Opus 5 * fix(signals): restore inputs= on the batched subscribe_to path On the subscribe_to path BatchSignalContext was built without inputs=, so b.inputs was {} and any signal declaring inputs=[...] raised KeyError on every call. sig/loss_debiased does exactly that: it failed 12,079 times in one five hour run -- once per step -- and because wrappered_fwd swallows per-signal exceptions nothing crashed, the column just silently never got values. We had already fixed this; taking dev src.py whole during the merge reverted it, since dev never carried the fix. Same class as the deque import and the label cache: dev has no equivalent, so a wholesale take drops it. Co-Authored-By: Claude Opus 5 * Remove optrace tracing from weightslab Drops the optrace module and every call site: 64 @traced decorators, 13 hit() markers and 5 imports across the data stores, the dataframe manager and the two services. Pure deletion -- 437 lines out, 0 in. Every hit() was verified to be a bare statement rather than an expression, so removing it cannot change a value, and removal was parenthesis-balanced because several spanned three lines. Each file is compiled after editing, which is what would catch a removal that left an empty block. The tracing was built to find where interactivity time went on a 100GB dataset. It has served that purpose: the O(change) view sync, the O(batch) NB_SEEN lookup and the flush accounting all came out of it. Co-Authored-By: Claude Opus 5 * ci: fix the code-quality and gRPC-test failures on this branch Three fixes, one per failing check. ruff F841, trainer_tools.process_sample: the positional unpack of _getitem_raw bound _res[1] to idx, which nothing reads -- the function returns sid. Dropped. ruff F401, examples/.../wl-video-generation/utils/data.py: unused `os` import. Pre-existing on dev and untouched by this branch; it only surfaces here because the lint step appends ./weightslab to the changed-file list, so ruff scans the whole package on any PR that touches it. Removing it is what unblocks the gate. AttributeError in tests/gRPC/test_grpc_user_actions.py: _fastUpdateInternals duck-types take_view_dirty and get_source_rows on the df manager. Both are new on this branch, so _FakeDFManager -- and any third-party manager -- raised AttributeError instead of taking the fallback. Guarded: a manager without dirty tracking cannot serve a delta, which is the same "structural change" case the method already falls back on, and the caller then runs _slowUpdateInternals exactly as before. Co-Authored-By: Claude Opus 5 Claude-Session: https://claude.ai/code/session_01RHV8zUd5aCtHtQourmwUJC * fix(agent): one shared OpenCode model for the studio, the CLI and the backend The backend answered on whatever model it resolved at start-up, and nothing could move it afterwards: a model picked in the studio was ignored, `agent model X` reported "Model switched to ", and `agent status` named a model that was no longer in use. OpenCode's own config (GET /config) is now the single shared choice, with: 1. OPENCODE_MODEL -- a hard pin, for automation. Never overridden; if the studio disagrees, `agent status` says so. 2. GET /config's model -- the live shared choice, RE-READ BEFORE EVERY TURN rather than latched at start-up. 3. agent_config.yaml's opencode_model -- a SEED: used when nothing has been chosen yet, and published so the studio shows it. It used to pin, so a run started after picking a model in the studio went back to the yaml value. 4. opencode/big-pickle -- the built-in default (was opencode/deepseek-v4-flash-free), also published. Whoever chooses last wins, and both surfaces follow. publish_model() writes PATCH /global/config (falling back to /config) and CONFIRMS by reading back: the workspace route answers 200 for a write it drops. An explicit switch (`agent model`, `agent init --model`, the RPC) is published rather than overwritten by the shared-config read, and current_model() re-resolves so `agent status` reports the model the NEXT query will use. The start-up banner now names the model actually in use and where it came from instead of printing "(server default)" whenever nothing was pinned. Co-Authored-By: Claude Opus 5 (1M context) * fix(cli): hand the experiment directory to runs started from another terminal `weightslab start` establishes the experiment directory and exports WEIGHTSLAB_ROOT_LOG_DIR -- into its OWN process only. A training run launched from a second terminal (or by `weightslab start example --seg`) is a different process tree and never saw it: it fell through to tempfile.mkdtemp(), so the run wrote reports/, notebooks/ and checkpoints into %TEMP%\tmpXXXXXXXX while the UI listed an empty reports/ from the directory it had established. That is the "right-click Generate report lists nothing, yet I generated reports" bug. weightslab/utils/active_experiment.py records the directory in a small per-user marker (~/.weightslab/active_experiment.json, WEIGHTSLAB_STATE_DIR to relocate), with two independent sections: `ui` (what `weightslab start` established) and `backend` (where training ACTUALLY resolved, whatever the route). Both writes are best-effort and every read validates the directory still exists. * root_log_dir resolution gains a step: explicit config > env > the recorded `ui` directory > temp dir. The temp-dir case now warns loudly that the UI will not find the run's files. * `weightslab start example` passes the recorded directory to the child, and says which directory it is using. Anything already set in the shell wins. * The UI's reports/notebooks/agent listings follow a LIVE backend's own recorded directory, so they stay right even when training was pointed elsewhere by a config file. Only a live one counts: the marker outlives the process that wrote it, and a finished run must not hijack the listing of a UI that was given its own directory. Co-Authored-By: Claude Opus 5 (1M context) * feat(ui): proxy the shared agent model same-origin The studio picker wrote OpenCode's shared model directly from the browser, which only works while OpenCode's --cors allowlist contains the page's exact origin. A LAN address, a tunnel hostname, or an `opencode serve` started by hand with no --cors all make that cross-origin PATCH fail its preflight, so the pick never reached OpenCode -- and therefore never reached weightslab's backend, which reads the same field to choose the model for its own queries. GET/POST /agent-server/model proxy it through this server instead: same-origin for the page, plain HTTP to OpenCode on the machine they share. The write goes to the global scope first and is confirmed by reading /config back, because the workspace route answers 200 for a write it drops. A transport failure or an error status is reported as ok:false rather than "nothing configured", so the page never shows its own default over the model actually in use. Co-Authored-By: Claude Opus 5 (1M context) * fix(data): stop a missing flag reading as a set one (phantom discards) On a detection/segmentation run, sample after sample greyed out as "discarded" in the studio as the model worked through the dataset -- while the dataframe said nothing was discarded. It came down to one line of Python semantics: bool(float("nan")) is True. The chain: 1. the trainer touches a sample -> its rows go dirty; 2. _fastUpdateInternals syncs the trainer-owned columns (signals*, last_seen, discarded, prediction, target) from the ledger into the served view. It collapsed the per-annotation rows with duplicated(keep="last"), keeping the LAST annotation row -- whose sample-level columns are NaN, because the real values live on the canonical row (annotation_id == 0, which is what the view itself is built from); 3. NaN therefore landed in the view's `discarded` (and in `prediction` / `target`, and `last_seen` went stale); 4. GetDataSamples served it as "1" if bool(value) else "0". Only annotation-expanded ledgers, only samples training had touched. Fixed in three places, plus the two siblings of the same bug found while auditing the function: * the differential sync now takes the canonical annotation_id == 0 row (falling back to the first occurrence), matching how the view is built; * is_set_flag / set_flag_mask replace bool() / astype(bool) wherever a nullable flag is read: the `discarded` rendering flag, the boolean tag:* columns in the metadata response -- where astype(bool) turned the NaN of every UNtagged sample into True, i.e. every sample wearing every tag -- and the histogram's per-(origin, discarded) split. They also read the strings "True"/"False" correctly, which a column that has been through the H5 store (categorical) can hold, and where bool("False") is True as well; * the sync's position lookup no longer searches the view's sample_id level, which raises InvalidIndexError as soon as one sample_id appears under two origins -- the very thing the view's (origin, sample_id) index exists to allow. It failed on every call there and silently fell back to the full rebuild. 13 tests reproduce the chain and each sibling; every one of them fails against the previous code. Co-Authored-By: Claude Opus 5 (1M context) * fix(cli): only a running `weightslab start` hands its experiment dir over The marker outlives the process that wrote it, so a directory recorded by a UI that had since exited redirected unrelated runs. It bit this repo's own suite: tests/gRPC/test_grpc_tag_operations.py resolved its root_log_dir into a previous session's experiment, found the segmentation example's config and checkpoints there, and errored in setUp with an unrelated config. * the handoff now reads live_ui_experiment_dir(), which requires the recording process to still be alive -- which is what "the UI is up over there, put this run in its experiment" actually means; * tests/conftest.py points WEIGHTSLAB_STATE_DIR at a throwaway directory for the whole session, so no test ever reads (or writes) the developer's own WeightsLab state, whatever a future one happens to resolve. Co-Authored-By: Claude Opus 5 (1M context) * fix(data): default `discarded` instead of leaving it NaN SampleStats.DEFAULTS documents `discarded` as False, directly under the comment "None are not accepted by PD H5 storage" -- so a NaN in it was already a broken contract, and it is where the phantom-discard bug started. The existing normalisation could not catch it: it only visits columns an upsert ADDS (`missing_cols`), and only when the incoming slice's dtype is already bool -- which it is not exactly when the slice carries missing values. So a sample registered without the flag, and every per-annotation row (sample-level values live on annotation 0), kept a NaN, and bool(NaN) is True. _fill_documented_flag_defaults() gives every column with a boolean default in SampleStats.DEFAULTS its default after each upsert. An isna().any() short-circuit per column means the common case touches no rows; a categorical column (what the H5 store hands back) is widened first, since fillna on a Categorical raises for a value outside its categories. Deliberately NOT applied to tag:* columns: for a boolean tag, NaN and False mean the same thing and NaN costs nothing, and for a categorical tag NaN means "unset", which is not a default at all. The read side (set_flag_mask) already treats both as not-set. Belt and braces with 4819483: the flag is defaulted at the source AND a NaN that reaches a reader anyway is read as not-set. Co-Authored-By: Claude Opus 5 (1M context) --------- Co-authored-by: Alexandru Rotaru Co-authored-by: Claude Opus 5 Co-authored-by: Guillaume --- agent_config.yaml | 17 +- docs/agent.rst | 78 ++- docs/perf/o_change_register.md | 170 ++++++ tests/conftest.py | 37 ++ tests/test_src_functions.py | 172 +++++- .../services/test_agent_opencode_provider.py | 117 ++++ .../test_data_service_discard_flag.py | 308 ++++++++++ tests/trainer/services/test_opencode_chat.py | 182 +++++- tests/ui/test_server_experiment_reports.py | 58 ++ tests/ui/test_server_shared_model.py | 184 ++++++ weightslab/backend/cli.py | 11 +- weightslab/backend/dataloader_interface.py | 48 ++ weightslab/backend/logger.py | 39 +- weightslab/cli.py | 46 ++ weightslab/data/dataframe_manager.py | 379 +++++++++++- weightslab/data/h5_array_store.py | 55 ++ weightslab/data/h5_dataframe_store.py | 136 ++++- .../PyTorch/wl-video-generation/utils/data.py | 1 - weightslab/src.py | 115 +++- weightslab/trainer/services/agent/agent.py | 173 +++++- .../trainer/services/agent/opencode_chat.py | 133 ++++- weightslab/trainer/services/data_service.py | 564 +++++++++++++++--- weightslab/trainer/trainer_tools.py | 13 +- weightslab/ui/server.py | 138 ++++- weightslab/utils/active_experiment.py | 340 +++++++++++ weightslab/utils/logs.py | 2 +- 26 files changed, 3326 insertions(+), 190 deletions(-) create mode 100644 docs/perf/o_change_register.md create mode 100644 tests/conftest.py create mode 100644 tests/trainer/services/test_data_service_discard_flag.py create mode 100644 tests/ui/test_server_shared_model.py create mode 100644 weightslab/utils/active_experiment.py diff --git a/agent_config.yaml b/agent_config.yaml index 3a3aa2f9..51b3c772 100644 --- a/agent_config.yaml +++ b/agent_config.yaml @@ -17,7 +17,16 @@ agent: opencode_url: http://127.0.0.1:4096 # OpenCode model, "providerID/modelID" (can also be set as env variable - # OPENCODE_MODEL). Empty string self-heals to whatever OpenCode's own - # config was last set to, or a configured provider default, falling back - # to the free-tier "opencode/deepseek-v4-flash-free" if neither resolves. - opencode_model: "opencode/deepseek-v4-flash-free" + # OPENCODE_MODEL). + # + # Left EMPTY on purpose. Empty means "follow OpenCode's own config" -- which + # is what the studio's model picker (.wl-ag-model) writes on every pick, via + # PUT /config -- so choosing a model in the UI also changes the model + # weightslab's own queries use, and it is re-checked before every turn. + # When OpenCode names no model either, the fallback is "opencode/big-pickle". + # + # A value here SEEDS the choice: it is used when nothing has been chosen in + # OpenCode's config yet (and is published there, so the studio shows it), but + # a model picked in the UI afterwards wins. Set OPENCODE_MODEL instead to + # PIN a model that nothing can override. + opencode_model: "" diff --git a/docs/agent.rst b/docs/agent.rst index 596555f4..9519418d 100644 --- a/docs/agent.rst +++ b/docs/agent.rst @@ -211,16 +211,71 @@ If ``OPENCODE_URL`` is set and reachable, the UI server adopts it directly instead of spawning a child; the backend SDK agent reads the same variable (``agent.py``'s ``_load_config``) — set it once and both sides talk to the one server, so a model you authenticate once is available everywhere. -``OPENCODE_MODEL`` (or ``agent_config.yaml``'s ``agent.opencode_model``) picks -the default model for the backend SDK agent, as an OpenCode -``providerID/modelID`` string (e.g. ``openrouter/anthropic/claude-opus-4.6``). -Leave it unset to fall back, in order, to: whatever model OpenCode's own -``/config`` was last set to (the model picker's own pick, e.g. from the -Weights Studio landing page), and otherwise the free-tier -``opencode/deepseek-v4-flash-free`` automatically — a provider's own -reported default used to be tried in between, but that could itself be an -arbitrary, non-text-reasoning model whenever any provider had credentials -configured, so it no longer overrides this. +Which model gets used, and how the two sides agree +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +**OpenCode's own config is the shared source of truth.** ``GET /config``'s +``model`` field is read by every client of that server -- the Weights Studio +model picker, the OpenCode CLI, and the backend SDK agent -- and written by the +picker (``PATCH /global/config``, ``global.config.update``) whenever you choose +a model. Neither side has to know the other exists; they meet in that one +field. + +The backend SDK agent resolves its model in this order: + +1. ``OPENCODE_MODEL`` -- a hard **pin**, for automation that must force a + model regardless of what anyone picked. The UI picker cannot move it; if the + two disagree, ``agent status`` logs which model this backend is actually + using and why. +2. ``GET /config``'s ``model`` -- the shared choice above. Re-read **before + every turn**, so switching models in the UI mid-run takes effect + immediately; it is not latched at start-up. +3. ``agent_config.yaml``'s ``agent.opencode_model`` -- a **seed**, not a pin: + which model to use when nothing has been chosen in OpenCode's config yet. + It is published there, so the studio shows it; once anything is chosen (in + the UI, by ``agent model``, or by the CLI), that choice wins and the seed + goes unused. Pinning the yaml value instead meant a run started *after* + picking a model in the studio quietly went back to the yaml one. +4. ``opencode/big-pickle``, the built-in fallback, when nothing above resolves. + Published too, so a backend that started before the UI hands the picker the + model it is itself using. + +A provider's own reported default used to be tried just before the built-in +fallback, but that could itself be an arbitrary, non-text-reasoning model +whenever any provider had credentials configured, so it no longer overrides it. + +Whoever chooses last wins, and both surfaces follow: picking in the UI moves +the backend's next query, and ``agent model `` from the +CLI moves the UI's picker. + +Either start order therefore converges on one model: + +.. code-block:: text + + Studio first: pick a model in the UI -> PATCH /global/config + -> weightslab start -> GET /config -> same model + (agent_config.yaml's seed is not used: something was chosen) + + weightslab first: nothing chosen anywhere + -> agent_config.yaml's opencode_model, else + opencode/big-pickle -- and published + -> Studio starts -> GET /config -> same model + +The start-up banner states which of the three won: + +.. code-block:: text + + Agent initialized from configuration C:\Users\you\wl_agent_config.yaml: + OpenCode URL=http://127.0.0.1:4096 + Model=opencode/muse-spark-1.3-contributor-free (from agent_config.yaml's opencode_model (nothing chosen yet), published to OpenCode's config so the studio shows it) + +Other values in the parentheses are ``pinned by OPENCODE_MODEL``, ``from +OpenCode's config, which the studio model picker writes``, ``chosen here and +published to OpenCode's config`` (an ``agent model`` switch), ``built-in +default, published to OpenCode's config for the studio``, and ``unresolved -- +OpenCode unreachable, retried on the first query``. It used to print ``Model=(server default)`` whenever nothing was +pinned, which read as "my pick was ignored" even when the first query would +have picked it up. Credentials and provider setup live in OpenCode itself, never in WeightsLab: @@ -335,7 +390,8 @@ configure ``agent_config.yaml`` and/or environment variables. # 3. agent_config.yaml agent: opencode_url: http://127.0.0.1:4096 - opencode_model: "" # empty = use OpenCode's own configured default + opencode_model: "" # empty = follow OpenCode's config (the UI picker); + # a value here PINS the model instead Then check it from the CLI: diff --git a/docs/perf/o_change_register.md b/docs/perf/o_change_register.md new file mode 100644 index 00000000..270f3417 --- /dev/null +++ b/docs/perf/o_change_register.md @@ -0,0 +1,170 @@ +# Interactivity for 100GB+ datasets — register of O(data) operations + +Goal: every per-step and per-request operation should cost **O(change)** or +**O(page)**, never **O(dataset)**. Today several do, so cost grows with the +dataset while the actual work stays constant. + +> **Baseline caveat.** The measurements below were taken against a stale +> in-place copy of weightslab (`~/weightslab_src`, 1.3.3+multiview), which +> differs from `dev` in 51 files. Two serving costs listed there are **already +> fixed on dev** (see §B). Serving numbers must be re-measured on this branch +> before any serving change is attributed an improvement. The storage findings +> (§A) were re-verified against dev and still hold. + +Reference measurements (UltraEdit, 3,959,093 rows, ~19 cols, A10G box): + +| | measured | +|---|---| +| bare-torch step (no WL) | 1,162 ms → 20.66 samples/s | +| with WL, no UI client | ~5.5 samples/s (**3.8× slower**) | +| with WL + image requests | ~1.5 samples/s (**13.8× slower**) | +| per-step signal write itself | **4 ms (0.34%)** — already fine | +| grid page latency (64 imgs) | p50 5.2 s, p95 9.6 s | + +The signal path is not the problem. Storage write-amplification and the view +rebuild are. + +--- + +## A. STORAGE — `data/h5_dataframe_store.py` + +`upsert()` **receives** only dirty rows but **implements** a full table +replacement. + +| line | operation | cost | +|---|---|---| +| 693 | `_create_backup()` — full file copy **before every upsert** | O(file) | +| 708 | `existing = store.select(key)` — read entire table | O(N) | +| ~768 | `pd.concat([existing, delta])` | O(N) | +| ~772 | `existing[~existing.index.duplicated()]` — dedupe all rows | O(N) | +| ~785 | `_decategorize_for_storage(existing)` | O(N) | +| 801 | `store.remove(key)` — drop table | O(N) | +| 804 | `store.append(..., data_columns=True)` — rewrite + index **every** column | O(N·cols) | +| 846–883 | same read/remove/rewrite in the column-delete path | O(N) | + +**Amplification:** ~5 KB of changed signals per flush → ~200 MB written, +roughly **40,000×**. At `ledger_flush_interval=3.0s` vs ~1.5 s steps, that is a +full-table rewrite about every 2 steps. + +Fix direction: append new rows; modify existing rows in place +(`select_as_coordinates` + `table.modify_rows`). Backup incrementally, not per +upsert. No `data_columns=True` — no `store.select()` in this file uses `where=`, +so those per-column indexes are built and never read. + +*(A previous attempt to narrow `data_columns` broke the write path entirely — +678 upsert failures, zero persisted data. Any change here needs a +write→read→assert-contents check, not just a timing check.)* + +## B. SERVING — `trainer/services/data_service.py` + +`_pull_into_all_data_view_df()` (line 938) runs several full-frame passes. + +**Already fixed on dev — do not re-report as wins:** +- the collapse no longer re-enters `get_combined_df()`; the pulled frame is + passed in, so the frame is copied once, not twice +- `array_proxy` no longer does a per-cell `.apply(convert_to_proxy)` (was + 1,660 ms at 4M rows) + +Remaining, to be **re-measured on this branch**: + +| line | operation | cost (measured @4M) | +|---|---|---| +| 946 | `get_combined_df()` → `dataframe_manager:2140 self._df.copy()` | 101 ms–1.2 s | +| — | `get_collapse_annotations_to_samples_df(df)` — groupby collapse | 6,326 ms* | +| — | `safe_reset_index(df)` | 1,912 ms | +| — | `set_index([origin, sample_id])` | O(N) | +| 3636 | `updated_df.reindex(target_order)` | 290 ms | + +Callers — each one is a full O(N) rebuild: lines **440, 851, 3593, 4605, 4626**, +reached from `GetDataSamples`, `GetMetaData`, `EditDataSample`, `GetDataSplits`. + +Held under `_update_lock`, which the trainer also needs → measured lock holds of +39–126 s and the 3.7× training penalty while browsing. + +**The collapse is provably a no-op when `annotation_id.max() == 0`** (UltraEdit +is exactly 1:1) yet still costs 6.3 s per rebuild. + +Fix direction: serve a page from the source frame by index (O(page)); rebuild +the full view only for genuinely global operations (histogram, global sort); +apply deltas rather than rebuilding; never hold the writer lock across a +rebuild — build off-lock and swap the reference. + +## C. OTHER FULL SCANS + +| location | note | +|---|---| +| `dataframe_manager:1875` `data_snapshot.iterrows()` | input is O(change), but row-wise Python per flush | +| `dataframe_manager:2400` `.apply(lambda …)` | per cell | +| `data_service:1200` `_compute_natural_sort_stats` | builds a list of one Series per row (4M objects). Gated off (`compute_natural_sort=False`) — latent | +| `data_service:538` PreviewCache | bounded by `WL_MAX_PREVIEW_CACHE_SIZE` — OK | + +## D. ALREADY O(change) — keep + +- `self._pending` dirty-row set (`dataframe_manager:95, 751, 761`) +- flush work set: `work = list(self._pending)` (`:1827`) +- `_origin_revisions` per-origin version counters (`:94, 1235`) + +The bookkeeping needed for differential updates already exists; the storage and +view layers just don't use it. + +\* measured on the stale copy; re-measure on dev. + +## Measurement protocol + +Fixed workload: **1,000 train samples** = 41 steps at batch 24. Every change is +reported as: + +1. wall-clock for the 41 steps, vs the bare-torch floor +2. bytes written to H5 for those steps +3. grid-page latency (64 images) and training throughput **while** serving +4. **ledger contents verified** — signal columns present, measured-row count + +(4) is not optional: a previous "10× win" was writes silently failing. + +--- + +# E. Triage — which call sites need a full reconstruction + +`_slowUpdateInternals()` rebuilds the whole view: `copy → collapse → reset_index +→ set_index → reindex`. It has 18 call sites, and almost none of them need +that. Most just want **fresh values for rows the trainer touched**, which is +`O(change)`. + +`_fastUpdateInternals()` applies only dirty rows, via a maintained +`sample_id → position` map (`_rebuild_view_pos_map`), and returns `False` — +falling back to the full rebuild — whenever it cannot safely apply: + +- no view yet, or no position map (first build) +- a dirty `sample_id` absent from the map (new rows ⇒ structural change) +- backlog > `max_dirty` (a rebuild is genuinely cheaper) + +So the worst case is today's behaviour, never wrong data. + +| site | routing | why | +|---|---|---| +| `_bg_view_refresh` | **fast** | exists purely to refresh values after a stale read — the textbook differential case, and the one that holds `_lock` against the trainer | +| `_process_get_data_samples` | **fast** | grid fetch needs current values, not a new frame | +| `_compute_custom_signals` | **fast** | writes new signal *values*; schema unchanged | +| `GetDataSplits` | **fast** | read-only summary | +| `EditDataSample` ×5 | **fast** | per-sample value edits | +| `EditDataSample` ×3 (`df.modify`, `df.drop_column`) | **full** | changes the schema — differential cannot add/remove columns | +| `ApplyDataQuery` `@reset`/`@clear` | **full** | clears `_is_filtered` to restore the full universe; a differential updates values but cannot restore *dropped rows* | +| `ApplyDataQuery` filter + agent paths | **full** (deferred) | a forced rebuild preserves `_is_filtered` (`:3691`), so swapping in a differential changes which rows the user sees. Rare, user-initiated, low perf value, high blast radius — not worth the risk until the filter semantics are pinned down | +| `_compute_natural_sort_stats` | **full** | gated off (`compute_natural_sort=False`); latent | +| `_manual_save_data_state` | **full** | explicit user save; correctness over speed | + +**Kill-switch:** `WL_FAST_VIEW=0` disables the differential and the position-map +build, reproducing prior behaviour exactly. This is what makes a like-for-like +A/B possible from a single tree. + +## Why the no-client benchmark cannot show this + +A 41-step run with no UI client attached records **0 rebuild events** — nothing +calls `_slowUpdateInternals` at all, so the fast path has nothing to improve and +correctly measures as no change. The rebuild cost only materialises when a +client is attached, which is the case that measured **3.7× slower** with p50 +grid latency of 5.2 s. + +The A/B is therefore run under load: baseline → under-load → recovery phases +within one training process (`t_imgload.py`), so each arm is normalised against +its own idle throughput. diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 00000000..32c86157 --- /dev/null +++ b/tests/conftest.py @@ -0,0 +1,37 @@ +"""Shared pytest setup. + +Keeps the suite out of the developer's own WeightsLab state. + +``weightslab.utils.active_experiment`` records the active experiment directory +in a real per-user file (``~/.weightslab/active_experiment.json``) so a +training run started in another terminal lands in the experiment the UI +established. That is deliberate at runtime -- and poison in tests: a run of +this suite while a ``weightslab start`` was up resolved its root_log_dir into +that live experiment, found the config and checkpoints of whatever was running +there, and failed in setUp with an unrelated config (seen for real: +tests/gRPC/test_grpc_tag_operations.py loading a segmentation example's +hyperparameters). + +So every test session gets its own throwaway state directory. Tests that +exercise the handoff itself set ``WEIGHTSLAB_STATE_DIR`` to their own temp +directory anyway; this only changes the default. +""" + +import os +import tempfile + +import pytest + + +@pytest.fixture(scope="session", autouse=True) +def _isolate_weightslab_state(): + previous = os.environ.get("WEIGHTSLAB_STATE_DIR") + with tempfile.TemporaryDirectory(prefix="wl-test-state-") as state_dir: + os.environ["WEIGHTSLAB_STATE_DIR"] = state_dir + try: + yield state_dir + finally: + if previous is None: + os.environ.pop("WEIGHTSLAB_STATE_DIR", None) + else: + os.environ["WEIGHTSLAB_STATE_DIR"] = previous diff --git a/tests/test_src_functions.py b/tests/test_src_functions.py index d99ed489..50c7a2b3 100644 --- a/tests/test_src_functions.py +++ b/tests/test_src_functions.py @@ -1,4 +1,6 @@ +import json import os +import shutil import tempfile import unittest import numpy as np @@ -6,6 +8,7 @@ import torch as th import weightslab.src as src +from weightslab.utils import active_experiment from unittest.mock import MagicMock, patch @@ -13,16 +16,27 @@ class TestResolveConfiguredRootLogDir(unittest.TestCase): - """root_log_dir resolution: explicit config > WEIGHTSLAB_ROOT_LOG_DIR > temp dir.""" + """root_log_dir resolution: explicit config > WEIGHTSLAB_ROOT_LOG_DIR > + the directory `weightslab start` recorded > temp dir.""" def setUp(self): self._env_prev = os.environ.get("WEIGHTSLAB_ROOT_LOG_DIR") + # The marker is a real per-user file; point it at a scratch directory so + # these tests never read (or write) the developer's own active run. + self._state_prev = os.environ.get("WEIGHTSLAB_STATE_DIR") + self._state_dir = tempfile.mkdtemp() + os.environ["WEIGHTSLAB_STATE_DIR"] = self._state_dir def tearDown(self): if self._env_prev is None: os.environ.pop("WEIGHTSLAB_ROOT_LOG_DIR", None) else: os.environ["WEIGHTSLAB_ROOT_LOG_DIR"] = self._env_prev + if self._state_prev is None: + os.environ.pop("WEIGHTSLAB_STATE_DIR", None) + else: + os.environ["WEIGHTSLAB_STATE_DIR"] = self._state_prev + shutil.rmtree(self._state_dir, ignore_errors=True) def test_explicit_config_value_wins_over_env(self): os.environ["WEIGHTSLAB_ROOT_LOG_DIR"] = "/env/dir" @@ -49,6 +63,162 @@ def test_falls_back_to_tempdir_when_neither_set(self): self.assertEqual(src._resolve_configured_root_log_dir(None), "/tmp/generated") mk.assert_called_once() + def test_a_recorded_directory_whose_ui_has_exited_is_not_adopted(self): + # The handoff means "the UI is up over there, join its experiment". A + # record left by a `weightslab start` that has since exited must not + # redirect an unrelated run -- it did, and this repo's own gRPC tests + # resolved into a previous session's experiment and loaded its config. + os.environ.pop("WEIGHTSLAB_ROOT_LOG_DIR", None) + with tempfile.TemporaryDirectory() as ui_dir: + active_experiment.record_ui_experiment(ui_dir) + state = active_experiment.read_state() + state["ui"][-1]["pid"] = 2 ** 31 - 1 # cannot be running + active_experiment.state_path().write_text(json.dumps(state), encoding="utf-8") + + with patch("weightslab.src.tempfile.mkdtemp", return_value="/tmp/generated") as mk: + self.assertEqual(src._resolve_configured_root_log_dir(None), "/tmp/generated") + mk.assert_called_once() + + def test_adopts_the_directory_weightslab_start_recorded(self): + # `weightslab start` exports WEIGHTSLAB_ROOT_LOG_DIR into its OWN + # process only. A training run in another terminal never saw it and + # went to a temp dir, so the UI listed an empty reports/ while the run + # wrote elsewhere. The recorded directory closes that gap. + os.environ.pop("WEIGHTSLAB_ROOT_LOG_DIR", None) + with tempfile.TemporaryDirectory() as ui_dir: + active_experiment.record_ui_experiment(ui_dir) + with patch("weightslab.src.tempfile.mkdtemp", return_value="/tmp/generated") as mk: + resolved = src._resolve_configured_root_log_dir(None) + mk.assert_not_called() + self.assertEqual(os.path.realpath(resolved), os.path.realpath(ui_dir)) + + def test_explicit_config_still_wins_over_the_recorded_directory(self): + with tempfile.TemporaryDirectory() as ui_dir: + active_experiment.record_ui_experiment(ui_dir) + self.assertEqual(src._resolve_configured_root_log_dir("/explicit/dir"), "/explicit/dir") + + def test_env_wins_over_the_recorded_directory(self): + with tempfile.TemporaryDirectory() as ui_dir, tempfile.TemporaryDirectory() as env_dir: + active_experiment.record_ui_experiment(ui_dir) + os.environ["WEIGHTSLAB_ROOT_LOG_DIR"] = env_dir + self.assertEqual(src._resolve_configured_root_log_dir(None), env_dir) + + def test_a_recorded_directory_that_no_longer_exists_is_ignored(self): + os.environ.pop("WEIGHTSLAB_ROOT_LOG_DIR", None) + gone = tempfile.mkdtemp() + active_experiment.record_ui_experiment(gone) + shutil.rmtree(gone, ignore_errors=True) + with patch("weightslab.src.tempfile.mkdtemp", return_value="/tmp/generated") as mk: + self.assertEqual(src._resolve_configured_root_log_dir(None), "/tmp/generated") + mk.assert_called_once() + + +class TestActiveExperimentMarker(unittest.TestCase): + """The cross-process handoff itself.""" + + def setUp(self): + self._state_prev = os.environ.get("WEIGHTSLAB_STATE_DIR") + self._state_dir = tempfile.mkdtemp() + os.environ["WEIGHTSLAB_STATE_DIR"] = self._state_dir + + def tearDown(self): + if self._state_prev is None: + os.environ.pop("WEIGHTSLAB_STATE_DIR", None) + else: + os.environ["WEIGHTSLAB_STATE_DIR"] = self._state_prev + shutil.rmtree(self._state_dir, ignore_errors=True) + + def test_two_uis_are_recorded_side_by_side(self): + # Two experiments at once (a cls UI and a seg UI) is supported; with one + # slot per side the second `weightslab start` erased the first. + with tempfile.TemporaryDirectory() as cls_dir, tempfile.TemporaryDirectory() as seg_dir: + active_experiment.record_ui_experiment(cls_dir, ui_port=8080, backend_port=50051) + # A second UI, standing in for another process. + state = active_experiment.read_state() + state["ui"].append({ + "root_log_dir": seg_dir, "pid": os.getpid(), + "ui_port": 8081, "backend_port": 50052, + }) + active_experiment.state_path().write_text(json.dumps(state), encoding="utf-8") + + recorded = {e["root_log_dir"] for e in active_experiment.entries("ui")} + self.assertEqual(len(recorded), 2) + + def test_two_live_uis_are_not_guessed_between(self): + with tempfile.TemporaryDirectory() as a_dir, tempfile.TemporaryDirectory() as b_dir: + active_experiment.record_ui_experiment(a_dir) + state = active_experiment.read_state() + state["ui"].append({"root_log_dir": b_dir, "pid": os.getpid()}) + active_experiment.state_path().write_text(json.dumps(state), encoding="utf-8") + + # Adopting either would put the run in the wrong experiment. + self.assertIsNone(active_experiment.live_ui_experiment_dir()) + + def test_a_backend_is_found_by_the_port_the_caller_talks_to(self): + with tempfile.TemporaryDirectory() as cls_dir, tempfile.TemporaryDirectory() as seg_dir: + state = {"backend": [ + {"root_log_dir": cls_dir, "pid": os.getpid(), "grpc_port": 50051}, + {"root_log_dir": seg_dir, "pid": os.getpid(), "grpc_port": 50052}, + ]} + active_experiment.state_path().parent.mkdir(parents=True, exist_ok=True) + active_experiment.state_path().write_text(json.dumps(state), encoding="utf-8") + + self.assertEqual( + os.path.realpath(active_experiment.live_backend_experiment_dir(50051)), + os.path.realpath(cls_dir)) + self.assertEqual( + os.path.realpath(active_experiment.live_backend_experiment_dir(50052)), + os.path.realpath(seg_dir)) + # No port, two candidates: no guess. + self.assertIsNone(active_experiment.live_backend_experiment_dir()) + # A port nobody serves: no guess either. + self.assertIsNone(active_experiment.live_backend_experiment_dir(50099)) + + def test_the_older_single_object_marker_is_still_readable(self): + with tempfile.TemporaryDirectory() as ui_dir: + active_experiment.state_path().parent.mkdir(parents=True, exist_ok=True) + active_experiment.state_path().write_text( + json.dumps({"ui": {"root_log_dir": ui_dir, "pid": os.getpid()}}), + encoding="utf-8") + self.assertEqual(os.path.realpath(active_experiment.ui_experiment_dir()), + os.path.realpath(ui_dir)) + self.assertEqual(os.path.realpath(active_experiment.live_ui_experiment_dir()), + os.path.realpath(ui_dir)) + + def test_ui_and_backend_entries_do_not_clobber_each_other(self): + with tempfile.TemporaryDirectory() as ui_dir, tempfile.TemporaryDirectory() as be_dir: + active_experiment.record_ui_experiment(ui_dir, ui_port=8080) + active_experiment.record_backend_experiment(be_dir) + self.assertEqual(os.path.realpath(active_experiment.ui_experiment_dir()), + os.path.realpath(ui_dir)) + self.assertEqual(os.path.realpath(active_experiment.backend_experiment_dir()), + os.path.realpath(be_dir)) + self.assertEqual(active_experiment.entries("ui")[-1]["ui_port"], 8080) + + def test_missing_marker_reads_as_nothing_recorded(self): + self.assertEqual(active_experiment.read_state(), {}) + self.assertIsNone(active_experiment.ui_experiment_dir()) + self.assertIsNone(active_experiment.backend_experiment_dir()) + + def test_a_corrupt_marker_is_ignored_rather_than_raising(self): + path = active_experiment.state_path() + path.parent.mkdir(parents=True, exist_ok=True) + path.write_text("{not json", encoding="utf-8") + self.assertEqual(active_experiment.read_state(), {}) + self.assertIsNone(active_experiment.ui_experiment_dir()) + + def test_recording_nothing_is_a_no_op(self): + self.assertIsNone(active_experiment.record_ui_experiment("")) + self.assertEqual(active_experiment.read_state(), {}) + + def test_clear_removes_the_marker(self): + with tempfile.TemporaryDirectory() as ui_dir: + active_experiment.record_ui_experiment(ui_dir) + self.assertTrue(active_experiment.state_path().exists()) + active_experiment.clear() + self.assertFalse(active_experiment.state_path().exists()) + active_experiment.clear() # idempotent + class TestSrcTagAndDiscardFunctions(unittest.TestCase): def setUp(self): diff --git a/tests/trainer/services/test_agent_opencode_provider.py b/tests/trainer/services/test_agent_opencode_provider.py index c2140950..868ba9c5 100644 --- a/tests/trainer/services/test_agent_opencode_provider.py +++ b/tests/trainer/services/test_agent_opencode_provider.py @@ -58,6 +58,118 @@ def _make_agent(df=None): return agent_mod, agent +def _agent_module(): + with mock.patch.dict(sys.modules, _install_agent_dependency_stubs(), clear=False): + return importlib.import_module("weightslab.trainer.services.agent.agent") + + +def _bare_agent(mod, model="opencode/ling-3.0-flash-fin-free"): + """An agent with just the OpenCode attributes _setup_providers touches. + + Sidesteps the full DataManipulationAgent construction (schema build, ctx, + handler registry) -- these tests are about the model plumbing only. + """ + agent = mod.DataManipulationAgent.__new__(mod.DataManipulationAgent) + agent.opencode_url = "http://127.0.0.1:4096" + agent.opencode_model = model + agent.opencode_workspace_dir = "." + agent._opencode_url_explicit = False + agent._opencode_model_explicit = False + agent._config_source_path = "(test)" + agent.preferred_provider = "opencode" + return agent + + +def _fake_chat(publish_ok=True, configured="opencode/ling-3.0-flash-fin-free"): + chat = MagicMock() + chat.base_url = "http://127.0.0.1:4096" + chat.model_is_explicit = False + chat.publish_model.return_value = publish_ok + # What the shared config would hand back if it were consulted. + chat.resolve_model.return_value = (configured, "opencode-config") + return chat + + +class TestModelSwitchAndReporting(unittest.TestCase): + """`agent model X` must actually switch to X, and `agent status` must name + the model the NEXT query will use. + + Regression: the shared-config re-read that lets the studio's model picker + move the backend also ran on the explicit switch path, overwriting the model + the caller had just asked for -- `agent model opencode/big-pickle` replied + "Model switched to opencode/ling-3.0-flash-fin-free". + """ + + def test_change_model_switches_and_publishes_the_choice(self): + mod = _agent_module() + agent = _bare_agent(mod) + chat = _fake_chat(publish_ok=True) + with mock.patch.object(mod, "OpenCodeChat", return_value=chat): + ok, message = agent.change_model("opencode/big-pickle") + + self.assertTrue(ok) + self.assertEqual(agent.opencode_model, "opencode/big-pickle") + self.assertIn("opencode/big-pickle", message) + self.assertNotIn("ling-3.0", message) + chat.publish_model.assert_called_once_with("opencode/big-pickle") + # The shared config must NOT be consulted on an explicit switch. + chat.resolve_model.assert_not_called() + # Published, so later turns keep following the config (the studio can + # still move it) rather than being pinned to this backend. + self.assertFalse(agent._opencode_model_explicit) + self.assertEqual(agent._opencode_model_source, "user-published") + + def test_change_model_pins_when_the_config_refuses_the_write(self): + mod = _agent_module() + agent = _bare_agent(mod) + chat = _fake_chat(publish_ok=False) + with mock.patch.object(mod, "OpenCodeChat", return_value=chat): + ok, message = agent.change_model("opencode/big-pickle") + + self.assertTrue(ok) + self.assertEqual(agent.opencode_model, "opencode/big-pickle") + self.assertTrue(agent._opencode_model_explicit) + self.assertTrue(chat.model_is_explicit) + self.assertEqual(agent._opencode_model_source, "user-pinned") + self.assertIn("pinned", message) + + def test_change_model_rejects_an_empty_model(self): + mod = _agent_module() + agent = _bare_agent(mod) + ok, message = agent.change_model(" ") + self.assertFalse(ok) + self.assertIn("empty", message.lower()) + + def test_startup_without_a_request_follows_the_shared_config(self): + mod = _agent_module() + agent = _bare_agent(mod, model="") + chat = _fake_chat(configured="opencode/muse-spark-1.3-contributor-free") + with mock.patch.object(mod, "OpenCodeChat", return_value=chat): + agent._setup_providers() + + chat.resolve_model.assert_called_once_with(publish_default=True) + self.assertEqual(agent.opencode_model, "opencode/muse-spark-1.3-contributor-free") + self.assertEqual(agent._opencode_model_source, "opencode-config") + + def test_current_model_reports_a_studio_pick_made_after_start_up(self): + mod = _agent_module() + agent = _bare_agent(mod, model="opencode/big-pickle") + chat = _fake_chat(configured="openrouter/openai/gpt-5-mini") + agent._opencode_chat = chat + + self.assertEqual(agent.current_model(), "openrouter/openai/gpt-5-mini") + self.assertEqual(agent.opencode_model, "openrouter/openai/gpt-5-mini") + + def test_current_model_keeps_the_last_known_model_when_the_server_is_down(self): + mod = _agent_module() + agent = _bare_agent(mod, model="opencode/big-pickle") + chat = _fake_chat() + chat.resolve_model.side_effect = OSError("refused") + agent._opencode_chat = chat + + self.assertEqual(agent.current_model(), "opencode/big-pickle") + + @unittest.skip("Not ready yet -- OpenCodeChat is still a stub, and the test suite needs to be reworked to support it") class TestOpenCodeConfigLoading(unittest.TestCase): def test_opencode_url_and_model_default(self): @@ -118,11 +230,16 @@ def test_opencode_chain_is_built(self): mock_cls.return_value.as_runnable.return_value = fake_runnable initialized = agent._setup_providers() + # seed_model carries agent_config.yaml's opencode_model, which SEEDS + # the shared choice instead of pinning it (a model picked in the studio + # wins over it); model/model_is_explicit carry OPENCODE_MODEL, which + # does pin. mock_cls.assert_called_once_with( agent.opencode_url, agent.opencode_model, workspace_dir=agent.opencode_workspace_dir, url_is_explicit=agent._opencode_url_explicit, model_is_explicit=agent._opencode_model_explicit, + seed_model=agent._opencode_model_seed, ) self.assertTrue(initialized) self.assertIs(agent.chain_opencode, fake_runnable) diff --git a/tests/trainer/services/test_data_service_discard_flag.py b/tests/trainer/services/test_data_service_discard_flag.py new file mode 100644 index 00000000..97da1fa5 --- /dev/null +++ b/tests/trainer/services/test_data_service_discard_flag.py @@ -0,0 +1,308 @@ +"""The 'samples discard themselves in the studio while training' regression. + +Reported for detection/segmentation runs: as the model worked through the +dataset, sample after sample greyed out in the UI, while the dataframe (checked +from the notebook) said nothing was discarded. + +The chain, reproduced below: + +1. the trainer touches a sample -> its rows go dirty; +2. `_fastUpdateInternals` syncs the columns the trainer mutates + (``signals*``, ``last_seen``, ``discarded``, ``prediction``, ``target``) + from the source into the view. It collapsed the source's per-annotation rows + with ``duplicated(keep="last")``, i.e. it kept the LAST annotation row -- + whose *sample-level* columns are NaN, because the real values live on the + canonical row (annotation_id == 0); +3. NaN therefore landed in the view's ``discarded``; +4. GetDataSamples reported that flag as ``"1" if bool(value) else "0"`` -- and + ``bool(float("nan"))`` is True in Python. + +So every sample the model had seen was served as discarded. Only for +annotation-expanded (detection/segmentation) ledgers, and only after training +touched the sample: exactly the reported shape. +""" + +import unittest + +import numpy as np +import pandas as pd + +from weightslab.data.sample_stats import SampleStatsEx +from weightslab.trainer.services.data_service import ( + DataService, + is_set_flag, + set_flag_mask, +) + + +SID = SampleStatsEx.SAMPLE_ID.value +ANNOT = SampleStatsEx.INSTANCE_ID.value +DISCARDED = SampleStatsEx.DISCARDED.value + + +def _source_rows(): + """A segmentation ledger: sample-level values on annotation 0 only.""" + return pd.DataFrame( + { + DISCARDED: [False, np.nan, np.nan, True, np.nan], + "last_seen": [11, 11, 11, 12, 12], + "prediction": ["a", None, None, "b", None], + }, + index=pd.MultiIndex.from_tuples( + [("0", 0), ("0", 1), ("0", 2), ("1", 0), ("1", 1)], + names=[SID, ANNOT], + ), + ) + + +class _FakeManager: + """Minimal dataframe manager: dirty tracking + source row lookup.""" + + def __init__(self, source, dirty): + self._df = source + self._dirty = list(dirty) + + def take_view_dirty(self, limit=None): + dirty, self._dirty = self._dirty, [] + return dirty + + def get_source_rows(self, sample_ids, columns=None): + wanted = [str(s) for s in sample_ids] + level = self._df.index.get_level_values(SID).astype(str) + rows = self._df[level.isin(wanted)] + return rows[columns] if columns else rows + + +class TestFastViewSyncKeepsSampleLevelValues(unittest.TestCase): + def _service(self, source, dirty, view): + service = DataService.__new__(DataService) + service._all_datasets_df = view + service._df_manager = _FakeManager(source, dirty) + return service + + def _collapsed_view(self, source): + """One row per sample, as the real view is built: annotation 0 wins.""" + base = source[source.index.get_level_values(ANNOT) == 0].droplevel(ANNOT) + return base.copy() + + def test_a_seen_sample_keeps_its_discarded_flag(self): + source = _source_rows() + view = self._collapsed_view(source) + service = self._service(source, ["0", "1"], view) + + self.assertTrue(service._fastUpdateInternals()) + + # Sample 0 is NOT discarded and must stay that way; sample 1 is. + self.assertIs(bool(view.loc["0", DISCARDED]), False) + self.assertIs(bool(view.loc["1", DISCARDED]), True) + # The give-away of the old behaviour: a NaN in a column that had a value. + self.assertFalse(view[DISCARDED].isna().any(), + "sample-level flags were overwritten with an instance row's NaN") + + def test_the_columns_the_trainer_owns_still_sync(self): + source = _source_rows() + view = self._collapsed_view(source) + source.loc[("0", 0), "last_seen"] = 99 + service = self._service(source, ["0"], view) + + self.assertTrue(service._fastUpdateInternals()) + self.assertEqual(view.loc["0", "last_seen"], 99) + + def test_falls_back_to_the_first_row_when_no_canonical_row_is_present(self): + # A slice of instance rows only (no annotation 0) must still sync + # something sane rather than raising or inventing a flag. + source = _source_rows().drop(index=("0", 0)) + view = self._collapsed_view(_source_rows()) + service = self._service(source, ["0"], view) + + self.assertTrue(service._fastUpdateInternals()) + self.assertTrue(pd.isna(view.loc["0", DISCARDED]) or view.loc["0", DISCARDED] is False) + + +class TestDiscardedFlagIsNaNSafe(unittest.TestCase): + """The second half: how a nullable flag becomes what the studio renders. + + is_set_flag / set_flag_mask are shared by the three places that read one: + GetDataSamples' `discarded` rendering flag, the boolean ``tag:*`` columns in + the metadata response, and the histogram's per-(origin, discarded) split. + """ + + @staticmethod + def _served_flag(value): + return "1" if is_set_flag(value) else "0" + + def test_missing_is_not_discarded(self): + # bool(float("nan")) is True -- the whole bug in one line. + self.assertEqual(self._served_flag(np.nan), "0") + self.assertEqual(self._served_flag(None), "0") + self.assertEqual(self._served_flag(pd.NA), "0") + + def test_real_values_still_come_through(self): + self.assertEqual(self._served_flag(True), "1") + self.assertEqual(self._served_flag(np.bool_(True)), "1") + self.assertEqual(self._served_flag(1), "1") + self.assertEqual(self._served_flag(False), "0") + self.assertEqual(self._served_flag(0), "0") + + + def test_a_string_flag_is_read_as_a_word_not_as_a_non_empty_string(self): + # bool("False") is True, and a boolean column that has been through the + # H5 store (categorical) can come back as these strings. + self.assertEqual(self._served_flag("False"), "0") + self.assertEqual(self._served_flag("false"), "0") + self.assertEqual(self._served_flag("0"), "0") + self.assertEqual(self._served_flag(""), "0") + self.assertEqual(self._served_flag("True"), "1") + self.assertEqual(self._served_flag("true"), "1") + self.assertEqual(self._served_flag("1"), "1") + + +class TestSetFlagMask(unittest.TestCase): + """The column-wide form, used for tags and the histogram split.""" + + def test_a_sparse_tag_column_marks_only_the_tagged_samples(self): + # A tag is set on a few samples; every other row is NaN. astype(bool) + # turned those into True -- every sample wore every tag. + column = pd.Series([True, np.nan, False, None, True], dtype=object) + self.assertEqual(set_flag_mask(column).tolist(), + [True, False, False, False, True]) + + def test_it_handles_a_categorical_column(self): + # The H5 store optimises tag:* and discarded to categorical dtype. + column = pd.Series([True, None, False], dtype=object).astype("category") + self.assertEqual(set_flag_mask(column).tolist(), [True, False, False]) + + def test_it_handles_a_float_column_of_zeros_and_nans(self): + column = pd.Series([1.0, np.nan, 0.0]) + self.assertEqual(set_flag_mask(column).tolist(), [True, False, False]) + + def test_an_absent_column_is_all_false(self): + self.assertEqual(set_flag_mask(None).tolist(), []) + + +class TestFastViewSyncAddressing(unittest.TestCase): + """How dirty source rows are matched to view rows. + + The view is indexed (origin, sample_id) precisely because one sample_id can + appear under two origins. Looking positions up in that non-unique level + raised InvalidIndexError, so the differential refresh failed on every call + and quietly fell back to the full rebuild. + """ + + def _service(self, source, dirty, view): + service = DataService.__new__(DataService) + service._all_datasets_df = view + service._df_manager = _FakeManager(source, dirty) + return service + + def _source(self): + return pd.DataFrame( + {DISCARDED: [False, np.nan], "last_seen": [7, 7]}, + index=pd.MultiIndex.from_tuples([("5", 0), ("5", 1)], names=[SID, ANNOT]), + ) + + def test_one_sample_id_under_two_origins_does_not_raise(self): + view = pd.DataFrame( + {DISCARDED: [False, False], "last_seen": [1, 2]}, + index=pd.MultiIndex.from_tuples([("train_loader", "5"), ("test_loader", "5")], + names=["origin", SID]), + ) + service = self._service(self._source(), ["5"], view) + + self.assertTrue(service._fastUpdateInternals()) + # The source can only speak per sample_id, so both rows take its values. + self.assertEqual(view["last_seen"].tolist(), [7, 7]) + self.assertFalse(view[DISCARDED].isna().any()) + + def test_a_dirty_sample_the_view_does_not_hold_forces_a_rebuild(self): + view = pd.DataFrame( + {DISCARDED: [False], "last_seen": [1]}, + index=pd.MultiIndex.from_tuples([("train_loader", "5")], names=["origin", SID]), + ) + source = pd.concat([self._source(), pd.DataFrame( + {DISCARDED: [False], "last_seen": [3]}, + index=pd.MultiIndex.from_tuples([("99", 0)], names=[SID, ANNOT]))]) + service = self._service(source, ["5", "99"], view) + + # 99 is new -> structural change -> only the full rebuild can add it. + self.assertFalse(service._fastUpdateInternals()) + + def test_nothing_to_do_when_no_dirty_sample_is_in_the_view(self): + view = pd.DataFrame( + {DISCARDED: [False], "last_seen": [1]}, + index=pd.MultiIndex.from_tuples([("train_loader", "7")], names=["origin", SID]), + ) + service = self._service(self._source(), ["5"], view) + self.assertTrue(service._fastUpdateInternals()) + self.assertEqual(view["last_seen"].tolist(), [1]) + + +class TestDocumentedFlagDefaults(unittest.TestCase): + """`discarded` should never be NaN in the first place. + + SampleStats.DEFAULTS documents it as False, right under the comment "None + are not accepted by PD H5 storage". It held NaN anyway: the existing + normalisation only visits columns an upsert ADDS, and only when the + incoming slice's dtype is already bool -- which it is not exactly when the + slice carries missing values. So a sample registered without the flag, and + every per-annotation row (sample-level values live on annotation 0), kept a + NaN, which `bool()` then read as True. + + Belt and braces with the read-side fix above: the flag is defaulted at the + source, AND a NaN that reaches a reader anyway is read as not-set. + """ + + def _manager(self, frame): + from weightslab.data.dataframe_manager import LedgeredDataFrameManager + manager = LedgeredDataFrameManager.__new__(LedgeredDataFrameManager) + manager._df = frame + return manager + + def test_missing_discarded_becomes_false(self): + frame = pd.DataFrame( + {DISCARDED: [True, np.nan, np.nan, False, np.nan]}, + index=pd.MultiIndex.from_tuples( + [("0", 0), ("0", 1), ("0", 2), ("1", 0), ("1", 1)], names=[SID, ANNOT]), + ) + manager = self._manager(frame) + manager._fill_documented_flag_defaults() + + self.assertEqual(frame[DISCARDED].tolist(), [True, False, False, False, False]) + self.assertFalse(frame[DISCARDED].isna().any()) + + def test_a_categorical_flag_column_is_widened_rather_than_raising(self): + # fillna on a Categorical raises unless the value is a known category, + # and the H5 store hands these columns back as categorical. + # NB: build the frame first, THEN cast -- handing a Series with its own + # RangeIndex to DataFrame(index=[...]) reindexes it to all-NaN. + frame = pd.DataFrame({DISCARDED: [True, None]}, index=pd.Index(["0", "1"], name=SID)) + frame[DISCARDED] = frame[DISCARDED].astype("category") + self.assertIsInstance(frame[DISCARDED].dtype, pd.CategoricalDtype) + manager = self._manager(frame) + manager._fill_documented_flag_defaults() + self.assertEqual(frame[DISCARDED].tolist(), [True, False]) + + def test_sparse_tag_columns_are_left_alone(self): + # NaN and False mean the same thing for a boolean tag, and NaN is + # cheaper; for a CATEGORICAL tag, NaN means "unset", not a default. + frame = pd.DataFrame({"tag:hard": [True, np.nan], DISCARDED: [False, np.nan]}, + index=pd.Index(["0", "1"], name=SID)) + manager = self._manager(frame) + manager._fill_documented_flag_defaults() + self.assertTrue(pd.isna(frame["tag:hard"].iloc[1])) + self.assertIs(bool(frame[DISCARDED].iloc[1]), False) + + def test_a_column_with_nothing_missing_is_not_rewritten(self): + frame = pd.DataFrame({DISCARDED: [True, False]}, index=pd.Index(["0", "1"], name=SID)) + before = frame[DISCARDED].dtype + manager = self._manager(frame) + manager._fill_documented_flag_defaults() + self.assertEqual(frame[DISCARDED].dtype, before) + + def test_an_empty_frame_is_a_no_op(self): + frame = pd.DataFrame() + self._manager(frame)._fill_documented_flag_defaults() # must not raise + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/trainer/services/test_opencode_chat.py b/tests/trainer/services/test_opencode_chat.py index 81468e21..13d8caf5 100644 --- a/tests/trainer/services/test_opencode_chat.py +++ b/tests/trainer/services/test_opencode_chat.py @@ -398,15 +398,37 @@ def test_explicit_model_is_never_auto_replaced(self): request_mock.assert_not_called() self.assertEqual(chat.model, "openrouter/anthropic/claude-opus-4.6") - def test_already_set_non_explicit_model_is_left_alone(self): - # Already resolved once (e.g. a prior call) -- don't re-resolve or - # re-request every single turn. + def test_already_set_non_explicit_model_follows_a_later_ui_pick(self): + # The studio's model picker writes the pick into OpenCode's own config + # (PUT /config). A backend that kept the model it resolved on its first + # turn went on answering with the old one for the rest of the run, so a + # NON-explicit model re-checks /config every turn and follows it. chat = OpenCodeChat("http://127.0.0.1:4096", model="openrouter/openai/gpt-5", model_is_explicit=False) + with mock.patch.object( + chat, "_request", + return_value=self._fake_response({"model": "opencode/big-pickle"}), + ) as request_mock: + chat._ensure_model_resolved() + request_mock.assert_called_once_with("/config") + self.assertEqual(chat.model, "opencode/big-pickle") + + def test_explicit_model_is_never_overridden_by_the_config(self): + # OPENCODE_MODEL / agent_config.yaml's opencode_model PIN the model: + # no request, and the UI picker cannot move it. + chat = OpenCodeChat("http://127.0.0.1:4096", model="openrouter/openai/gpt-5", model_is_explicit=True) with mock.patch.object(chat, "_request") as request_mock: chat._ensure_model_resolved() request_mock.assert_not_called() self.assertEqual(chat.model, "openrouter/openai/gpt-5") + def test_a_failed_config_read_keeps_the_model_already_resolved(self): + # One local request failing must not drop a working model for the + # fallback mid-run. + chat = OpenCodeChat("http://127.0.0.1:4096", model="anthropic/claude-haiku-4.5", model_is_explicit=False) + with mock.patch.object(chat, "_request", side_effect=OSError("refused")): + chat._ensure_model_resolved() + self.assertEqual(chat.model, "anthropic/claude-haiku-4.5") + def test_unset_model_resolves_from_config_own_model_field(self): chat = OpenCodeChat("http://127.0.0.1:4096", model=None, model_is_explicit=False) with mock.patch.object(chat, "_request", return_value=self._fake_response({"model": "anthropic/claude-haiku-4.5"})) as request_mock: @@ -425,13 +447,13 @@ def test_falls_back_to_the_hardcoded_default_when_config_has_no_model(self): with mock.patch.object(chat, "_request", return_value=self._fake_response({})) as request_mock: chat._ensure_model_resolved() request_mock.assert_called_once_with("/config") - self.assertEqual(chat.model, "opencode/deepseek-v4-flash-free") + self.assertEqual(chat.model, "opencode/big-pickle") def test_no_resolvable_model_falls_back_to_the_hardcoded_default(self): chat = OpenCodeChat("http://127.0.0.1:4096", model=None, model_is_explicit=False) with mock.patch.object(chat, "_request", side_effect=OSError("refused")): chat._ensure_model_resolved() # must not raise - self.assertEqual(chat.model, "opencode/deepseek-v4-flash-free") + self.assertEqual(chat.model, "opencode/big-pickle") def test_config_field_that_is_not_provider_slash_model_falls_through(self): # A malformed/unexpected `model` field (missing the "/") is treated @@ -441,7 +463,155 @@ def test_config_field_that_is_not_provider_slash_model_falls_through(self): with mock.patch.object(chat, "_request", return_value=self._fake_response({"model": "not-a-provider-model-pair"})) as request_mock: chat._ensure_model_resolved() request_mock.assert_called_once_with("/config") - self.assertEqual(chat.model, "opencode/deepseek-v4-flash-free") + self.assertEqual(chat.model, "opencode/big-pickle") + + + +class PublishModelTests(unittest.TestCase): + """publish_model / resolve_model -- how the two sides converge on one model. + + Verified against a live server before these were written: PATCH /config + (workspace scope) answers 200 and echoes the value back but does NOT change + what GET /config reports, while PATCH /global/config does. GET /config is + what both the studio picker and this backend read, so the write has to go + to the global scope, and be CONFIRMED rather than trusted. + """ + + @staticmethod + def _resp(payload): + return mock.MagicMock( + __enter__=mock.MagicMock( + return_value=mock.MagicMock(read=lambda: json.dumps(payload).encode())), + __exit__=mock.MagicMock(return_value=False), + ) + + def test_publish_writes_the_global_scope_first(self): + chat = OpenCodeChat("http://127.0.0.1:4096", model="opencode/big-pickle", + model_is_explicit=False) + calls = [] + + def fake_request(path, method="GET", body=None, headers=None): + calls.append((path, method, body)) + if method == "GET": + return self._resp({"model": "opencode/big-pickle"}) + return self._resp({}) + + with mock.patch.object(chat, "_request", side_effect=fake_request): + self.assertTrue(chat.publish_model()) + self.assertEqual(calls[0], ("/global/config", "PATCH", {"model": "opencode/big-pickle"})) + # Confirmed by reading the effective config back, not by the echo. + self.assertIn(("/config", "GET", None), calls) + + def test_publish_falls_back_to_the_workspace_route(self): + chat = OpenCodeChat("http://127.0.0.1:4096", model="opencode/big-pickle", + model_is_explicit=False) + seen = [] + + def fake_request(path, method="GET", body=None, headers=None): + seen.append((path, method)) + if path == "/global/config": + raise OSError("no such route") + if method == "GET": + return self._resp({"model": "opencode/big-pickle"}) + return self._resp({}) + + with mock.patch.object(chat, "_request", side_effect=fake_request): + self.assertTrue(chat.publish_model()) + self.assertIn(("/config", "PATCH"), seen) + + def test_publish_reports_failure_when_the_write_does_not_stick(self): + # Exactly the live failure mode: 200 back, effective config unchanged. + chat = OpenCodeChat("http://127.0.0.1:4096", model="opencode/big-pickle", + model_is_explicit=False) + + def fake_request(path, method="GET", body=None, headers=None): + if method == "GET": + return self._resp({}) # no model -- nothing was stored + return self._resp({"model": "opencode/big-pickle"}) # echo only + + with mock.patch.object(chat, "_request", side_effect=fake_request): + self.assertFalse(chat.publish_model()) + + def test_resolve_publishes_the_fallback_so_the_studio_adopts_it(self): + # weightslab started FIRST: nothing chosen anywhere, so the built-in + # default is resolved AND published for the UI to read. + chat = OpenCodeChat("http://127.0.0.1:4096", model="", model_is_explicit=False) + with mock.patch.object(chat, "_configured_model", side_effect=[None, "opencode/big-pickle"]), mock.patch.object(chat, "_request", return_value=self._resp({})), mock.patch.object(chat, "_ensure_reachable"): + model, source = chat.resolve_model(publish_default=True) + self.assertEqual(model, "opencode/big-pickle") + self.assertEqual(source, "default-published") + + def test_resolve_adopts_the_studio_pick_without_publishing(self): + # studio started FIRST: its pick is already in OpenCode's config, so + # follow it and write nothing. + chat = OpenCodeChat("http://127.0.0.1:4096", model="", model_is_explicit=False) + with mock.patch.object(chat, "_configured_model", return_value="openrouter/openai/gpt-5-mini"), mock.patch.object(chat, "publish_model") as publish, mock.patch.object(chat, "_ensure_reachable"): + model, source = chat.resolve_model(publish_default=True) + self.assertEqual((model, source), ("openrouter/openai/gpt-5-mini", "opencode-config")) + publish.assert_not_called() + + def test_the_studio_pick_wins_over_a_yaml_seed(self): + # The reported bug: pick a model in the UI, then start a run -- + # agent_config.yaml's opencode_model pinned the backend back to its own + # value, ignoring the pick. A yaml model is a SEED, not a pin. + chat = OpenCodeChat("http://127.0.0.1:4096", model="", model_is_explicit=False, + seed_model="opencode/muse-spark-1.3-contributor-free") + with mock.patch.object(chat, "_configured_model", return_value="openrouter/openai/gpt-5-mini"), mock.patch.object(chat, "publish_model") as publish, mock.patch.object(chat, "_ensure_reachable"): + model, source = chat.resolve_model(publish_default=True) + self.assertEqual((model, source), ("openrouter/openai/gpt-5-mini", "opencode-config")) + publish.assert_not_called() + + def test_the_yaml_seed_is_used_and_published_when_nothing_is_chosen(self): + chat = OpenCodeChat("http://127.0.0.1:4096", model="", model_is_explicit=False, + seed_model="opencode/muse-spark-1.3-contributor-free") + with mock.patch.object(chat, "_configured_model", return_value=None), mock.patch.object(chat, "publish_model", return_value=True) as publish, mock.patch.object(chat, "_ensure_reachable"): + model, source = chat.resolve_model(publish_default=True) + self.assertEqual( + (model, source), + ("opencode/muse-spark-1.3-contributor-free", "config-seed-published")) + publish.assert_called_once() + + def test_the_yaml_seed_beats_the_builtin_default(self): + chat = OpenCodeChat("http://127.0.0.1:4096", model="", model_is_explicit=False, + seed_model="openrouter/openai/gpt-5-mini") + with mock.patch.object(chat, "_configured_model", return_value=None), mock.patch.object(chat, "publish_model", return_value=False), mock.patch.object(chat, "_ensure_reachable"): + model, _ = chat.resolve_model(publish_default=True) + self.assertEqual(model, "openrouter/openai/gpt-5-mini") + + def test_the_env_pin_still_beats_the_studio_pick(self): + # OPENCODE_MODEL is per-process and deliberate: automation must be able + # to force a model regardless of what anyone picked in the UI. + chat = OpenCodeChat("http://127.0.0.1:4096", model="openrouter/pinned/model", + model_is_explicit=True, seed_model="opencode/seed") + with mock.patch.object(chat, "_configured_model", return_value="openrouter/openai/gpt-5-mini"), mock.patch.object(chat, "publish_model", return_value=True), mock.patch.object(chat, "_ensure_reachable"): + model, source = chat.resolve_model(publish_default=True) + self.assertEqual((model, source), ("openrouter/pinned/model", "pinned-published")) + + def test_resolve_publishes_a_pinned_model_so_the_studio_shows_it(self): + # A model pinned in agent_config.yaml / OPENCODE_MODEL is a deliberate + # choice too. It used to stay invisible to the studio, which then + # displayed an unrelated default while every backend query ran on the + # pinned model. + chat = OpenCodeChat("http://127.0.0.1:4096", model="openrouter/pinned/model", + model_is_explicit=True) + with mock.patch.object(chat, "publish_model", return_value=True) as publish, mock.patch.object(chat, "_ensure_reachable"): + model, source = chat.resolve_model(publish_default=True) + self.assertEqual((model, source), ("openrouter/pinned/model", "pinned-published")) + publish.assert_called_once() + + def test_a_pinned_model_stays_pinned_when_it_cannot_be_published(self): + chat = OpenCodeChat("http://127.0.0.1:4096", model="openrouter/pinned/model", + model_is_explicit=True) + with mock.patch.object(chat, "publish_model", return_value=False), mock.patch.object(chat, "_ensure_reachable"): + model, source = chat.resolve_model(publish_default=True) + self.assertEqual((model, source), ("openrouter/pinned/model", "pinned")) + + def test_a_model_read_from_the_config_is_never_republished(self): + chat = OpenCodeChat("http://127.0.0.1:4096", model="", model_is_explicit=False) + with mock.patch.object(chat, "_configured_model", return_value="opencode/big-pickle"), mock.patch.object(chat, "publish_model") as publish, mock.patch.object(chat, "_ensure_reachable"): + model, source = chat.resolve_model(publish_default=True) + self.assertEqual((model, source), ("opencode/big-pickle", "opencode-config")) + publish.assert_not_called() if __name__ == "__main__": diff --git a/tests/ui/test_server_experiment_reports.py b/tests/ui/test_server_experiment_reports.py index bd0c0e7a..b1cf36ee 100644 --- a/tests/ui/test_server_experiment_reports.py +++ b/tests/ui/test_server_experiment_reports.py @@ -17,6 +17,7 @@ import json import os +import shutil import tempfile import threading import time @@ -25,6 +26,7 @@ import urllib.request from weightslab.ui import server as ui_server +from weightslab.utils import active_experiment class _ServerTestCase(unittest.TestCase): @@ -33,6 +35,14 @@ class _ServerTestCase(unittest.TestCase): def setUp(self): self.tmp = tempfile.mkdtemp() + # The active-experiment marker is a real per-user file, and a LIVE + # backend recorded in it redirects these listings on purpose (see + # _experiment_dir_path). Point it at a scratch directory so the tests + # exercise the explicit experiment_dir below instead of whatever run + # the developer happens to have going. + self._state_prev = os.environ.get("WEIGHTSLAB_STATE_DIR") + self._state_dir = tempfile.mkdtemp() + os.environ["WEIGHTSLAB_STATE_DIR"] = self._state_dir self.httpd = ui_server.serve_ui( ui_host="127.0.0.1", ui_port=0, backend_host="localhost", backend_port=50051, @@ -47,6 +57,11 @@ def setUp(self): def tearDown(self): self.httpd.shutdown() self.thread.join(timeout=5) + if self._state_prev is None: + os.environ.pop("WEIGHTSLAB_STATE_DIR", None) + else: + os.environ["WEIGHTSLAB_STATE_DIR"] = self._state_prev + shutil.rmtree(self._state_dir, ignore_errors=True) def _get(self, path): return urllib.request.urlopen(f"http://127.0.0.1:{self.port}{path}", timeout=5) @@ -96,6 +111,49 @@ def test_entries_include_name_path_and_modified_at(self): self.assertIsInstance(entry["modified_at"], (int, float)) +class TestReportsFollowTheRunningBackend(_ServerTestCase): + """The listing follows where the RUNNING backend actually writes. + + Regression: `weightslab start` and a training run started in another + terminal resolved different directories (the run fell through to %TEMP%), + so right-clicking "Generate report" listed nothing even though reports had + just been generated. + """ + + def _write_report_in(self, directory, name): + reports_dir = os.path.join(directory, "reports") + os.makedirs(reports_dir, exist_ok=True) + with open(os.path.join(reports_dir, name), "w", encoding="utf-8") as f: + f.write("") + + def test_a_live_backend_directory_is_listed_instead_of_the_uis_own(self): + backend_dir = tempfile.mkdtemp() + self.addCleanup(shutil.rmtree, backend_dir, ignore_errors=True) + self._write_report_in(backend_dir, "experiment_report_20260909_193124.html") + # This test process stands in for the running backend. + active_experiment.record_backend_experiment(backend_dir) + + with self._get("/experiment-report/list") as r: + data = json.loads(r.read().decode()) + self.assertEqual([e["name"] for e in data["reports"]], + ["experiment_report_20260909_193124.html"]) + + def test_a_finished_backend_does_not_hijack_the_listing(self): + backend_dir = tempfile.mkdtemp() + self.addCleanup(shutil.rmtree, backend_dir, ignore_errors=True) + self._write_report_in(backend_dir, "stale.html") + active_experiment.record_backend_experiment(backend_dir) + # Rewrite the record with a pid that cannot be running. + state = active_experiment.read_state() + state["backend"][-1]["pid"] = 2 ** 31 - 1 + active_experiment.state_path().write_text(json.dumps(state), encoding="utf-8") + + self._write_report("mine.html") + with self._get("/experiment-report/list") as r: + data = json.loads(r.read().decode()) + self.assertEqual([e["name"] for e in data["reports"]], ["mine.html"]) + + class TestServeExperimentReport(_ServerTestCase): def test_serves_html_content_with_correct_content_type(self): diff --git a/tests/ui/test_server_shared_model.py b/tests/ui/test_server_shared_model.py new file mode 100644 index 00000000..58ffa1f1 --- /dev/null +++ b/tests/ui/test_server_shared_model.py @@ -0,0 +1,184 @@ +"""Tests for the UI server's same-origin shared-model endpoints: + +- GET /agent-server/model -- the model OpenCode's config names, or null. +- POST /agent-server/model -- set it (global scope, confirmed by reading back). + +Why they exist: the browser CAN call OpenCode directly, but only while +OpenCode's ``--cors`` allowlist contains the page's exact origin. A LAN +address, a tunnel hostname, or an ``opencode serve`` started by hand with no +``--cors`` all make that cross-origin PATCH fail its preflight, so a model +picked in the studio silently never reached OpenCode -- and therefore never +reached weightslab's backend, which reads that same field to choose the model +for its own queries. Proxying through this server is same-origin: no +preflight, no allowlist. + +A fake OpenCode stands in for the real server, reproducing the two behaviours +confirmed live: PATCH /global/config sticks, PATCH /config answers 200 and +echoes the value back without changing what GET /config reports. +""" + +import json +import threading +import time +import unittest +import unittest.mock +import urllib.error +import urllib.request +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer + +from weightslab.ui import server as ui_server + + +class _FakeOpencode(BaseHTTPRequestHandler): + """model lives in the class so every request sees the same value.""" + + model = None + workspace_patch_sticks = False + reachable = True + + def log_message(self, *args): # silence + pass + + def _json(self, status, payload): + body = json.dumps(payload).encode() + self.send_response(status) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(body))) + self.end_headers() + self.wfile.write(body) + + def do_GET(self): + if not type(self).reachable: + self._json(500, {"error": "down"}) + return + if self.path == "/config": + self._json(200, {"model": type(self).model} if type(self).model else {}) + return + self._json(404, {}) + + def do_PATCH(self): + length = int(self.headers.get("Content-Length", "0") or 0) + body = json.loads(self.rfile.read(length).decode() or "{}") if length else {} + model = body.get("model") + if self.path == "/global/config": + type(self).model = model + self._json(200, {"model": model}) + return + if self.path == "/config": + # Answers 200 and echoes the value, but only *stores* it when the + # server actually honours workspace scope -- the live one does not. + if type(self).workspace_patch_sticks: + type(self).model = model + self._json(200, {"model": model}) + return + self._json(404, {}) + + +class TestSharedModelEndpoints(unittest.TestCase): + def setUp(self): + _FakeOpencode.model = None + _FakeOpencode.workspace_patch_sticks = False + _FakeOpencode.reachable = True + + self.oc = ThreadingHTTPServer(("127.0.0.1", 0), _FakeOpencode) + self.oc_thread = threading.Thread(target=self.oc.serve_forever, daemon=True) + self.oc_thread.start() + oc_url = f"http://127.0.0.1:{self.oc.server_address[1]}" + + # The UI server resolves OpenCode from the running session, then from + # OPENCODE_URL -- point it at the fake. + self._prev_url = ui_server.os.environ.get("OPENCODE_URL") + ui_server.os.environ["OPENCODE_URL"] = oc_url + + self.httpd = ui_server.serve_ui( + ui_host="127.0.0.1", ui_port=0, + backend_host="localhost", backend_port=50051, + open_browser=False, block=False, + ) + self.port = self.httpd.server_address[1] + self.thread = threading.Thread(target=self.httpd.serve_forever, daemon=True) + self.thread.start() + time.sleep(0.1) + + def tearDown(self): + self.httpd.shutdown() + self.thread.join(timeout=5) + self.oc.shutdown() + self.oc_thread.join(timeout=5) + if self._prev_url is None: + ui_server.os.environ.pop("OPENCODE_URL", None) + else: + ui_server.os.environ["OPENCODE_URL"] = self._prev_url + + # -- helpers --------------------------------------------------------- + def _get(self): + with urllib.request.urlopen( + f"http://127.0.0.1:{self.port}/agent-server/model", timeout=5) as r: + return json.loads(r.read().decode()) + + def _post(self, model): + req = urllib.request.Request( + f"http://127.0.0.1:{self.port}/agent-server/model", method="POST", + data=json.dumps({"model": model}).encode(), + headers={"Content-Type": "application/json"}) + try: + with urllib.request.urlopen(req, timeout=5) as r: + return r.status, json.loads(r.read().decode()) + except urllib.error.HTTPError as exc: + return exc.code, json.loads(exc.read().decode() or "{}") + + # -- tests ----------------------------------------------------------- + def test_get_reports_nothing_when_no_model_is_configured(self): + self.assertEqual(self._get(), {"ok": True, "model": None}) + + def test_get_reports_the_configured_model(self): + _FakeOpencode.model = "opencode/big-pickle" + self.assertEqual(self._get(), {"ok": True, "model": "opencode/big-pickle"}) + + def test_get_ignores_a_malformed_model_field(self): + _FakeOpencode.model = "not-a-provider-model-pair" + self.assertEqual(self._get(), {"ok": True, "model": None}) + + def test_post_writes_the_global_scope_and_confirms_it(self): + status, payload = self._post("openrouter/openai/gpt-5-mini") + self.assertEqual(status, 200) + self.assertTrue(payload["ok"]) + self.assertEqual(payload["model"], "openrouter/openai/gpt-5-mini") + self.assertEqual(payload["via"], "/global/config") + # ...and it really is what OpenCode now reports. + self.assertEqual(_FakeOpencode.model, "openrouter/openai/gpt-5-mini") + self.assertEqual(self._get()["model"], "openrouter/openai/gpt-5-mini") + + def test_post_rejects_a_model_without_a_provider(self): + status, payload = self._post("big-pickle") + self.assertEqual(status, 400) + self.assertFalse(payload["ok"]) + + def test_post_reports_failure_when_the_write_does_not_stick(self): + # Both routes answer 200 but nothing is stored -- exactly the live + # workspace-scope behaviour, generalised. + class _EchoOnly(_FakeOpencode): + pass + + def do_PATCH(self): # noqa: N802 -- HTTP handler naming + length = int(self.headers.get("Content-Length", "0") or 0) + if length: + self.rfile.read(length) + self._json(200, {"model": "whatever"}) + + with unittest.mock.patch.object(_FakeOpencode, "do_PATCH", do_PATCH): + status, payload = self._post("opencode/big-pickle") + self.assertEqual(status, 200) + self.assertFalse(payload["ok"]) + self.assertIn("did not accept", payload["error"]) + + def test_get_says_so_when_opencode_is_unreachable(self): + _FakeOpencode.reachable = False + payload = self._get() + self.assertFalse(payload["ok"]) + self.assertIsNone(payload["model"]) + self.assertIn("not reachable", payload["error"]) + + +if __name__ == "__main__": + unittest.main() diff --git a/weightslab/backend/cli.py b/weightslab/backend/cli.py index c3d8a853..11ccbeb7 100644 --- a/weightslab/backend/cli.py +++ b/weightslab/backend/cli.py @@ -801,12 +801,21 @@ def _handle_command(cmd: str) -> Any: available = bool(agent.is_available()) except Exception: available = False + # current_model() re-reads OpenCode's shared config, so a model + # picked in the studio after this backend started is reported + # here instead of the model resolved at start-up. + model = getattr(agent, 'opencode_model', None) + try: + if hasattr(agent, 'current_model'): + model = agent.current_model() or model + except Exception: + pass return { 'ok': True, 'available': available, 'preferred_provider': getattr(agent, 'preferred_provider', None), 'opencode_url': getattr(agent, 'opencode_url', None), - 'opencode_model': getattr(agent, 'opencode_model', None), + 'opencode_model': model, 'message': 'Agent available. Ready to help you.' if available else 'Agent not configured. Use agent init.', } diff --git a/weightslab/backend/dataloader_interface.py b/weightslab/backend/dataloader_interface.py index 9d796a97..19f7f6b4 100644 --- a/weightslab/backend/dataloader_interface.py +++ b/weightslab/backend/dataloader_interface.py @@ -43,6 +43,44 @@ _DENY_LIST_REFRESH_INTERVAL = 32 +def _close_inherited_h5_fds(worker_id: int = 0) -> None: + """DataLoader worker_init_fn: drop HDF5 handles inherited from the parent. + + torch forks workers, so every fd the parent had open at fork time is + duplicated into the child -- including the ledger store. The child never + uses them, but their mere existence makes HDF5 refuse the parent's + read-write open, which silently kills ledger persistence. + + Closes rather than just dropping the Python object: the fd is what holds + the file, and the child has no Python-level reference to it at all. + """ + import os + try: + fd_dir = "/proc/self/fd" + for entry in os.listdir(fd_dir): + try: + target = os.readlink(os.path.join(fd_dir, entry)) + except OSError: + continue + if target.endswith(".h5") or target.endswith(".h5.lock"): + try: + os.close(int(entry)) + except OSError: + pass + except Exception: + # Never let cleanup break a worker: a leaked handle degrades + # persistence, a raising worker_init_fn kills the run. + pass + + +def _with_worker_init(kwargs: dict, num_workers: int) -> dict: + """Attach the fd cleanup unless the caller supplied its own init.""" + if num_workers and not kwargs.get("worker_init_fn"): + kwargs = dict(kwargs) + kwargs["worker_init_fn"] = _close_inherited_h5_fds + return kwargs + + def _resolve_safe_num_workers(dataset: Any, num_workers: int, loader_name: Optional[str] = None) -> int: """Clamp worker count for datasets that cannot be pickled by Windows spawn.""" try: @@ -177,6 +215,12 @@ def _get_deny_list_revision(self) -> Optional[tuple[str, int]]: try: origin = self._get_current_origin() df_manager = get_dataframe() + if origin and df_manager is not None and hasattr(df_manager, "get_discard_revision"): + # Deliberately NOT get_origin_revision: that moves on every + # per-sample signal write, i.e. every training step, so the + # cache below could never hit and __len__ rescanned 3.96M rows + # per batch. Discard state is what the deny-list depends on. + return ("discard", int(df_manager.get_discard_revision(origin))) if origin and df_manager is not None and hasattr(df_manager, "get_origin_revision"): return ("origin", int(df_manager.get_origin_revision(origin))) except Exception: @@ -514,6 +558,7 @@ def __init__( self.tracked_dataset, batch_sampler=batch_sampler, num_workers=num_workers, + worker_init_fn=_close_inherited_h5_fds, pin_memory=pin_memory, collate_fn=collate_fn, persistent_workers=self._should_persist_workers(num_workers), @@ -1134,6 +1179,7 @@ def restore_iteration_state(self, state: dict) -> None: self.tracked_dataset, batch_sampler=sampler, num_workers=num_workers, + worker_init_fn=_close_inherited_h5_fds, pin_memory=pin_memory, collate_fn=collate_fn, # Ensure no conflicting args are passed alongside batch_sampler @@ -1220,6 +1266,7 @@ def set_batch_size(self, new_batch_size: int) -> None: self.tracked_dataset, batch_sampler=sampler, num_workers=num_workers, + worker_init_fn=_close_inherited_h5_fds, pin_memory=pin_memory, collate_fn=collate_fn, persistent_workers=self._should_persist_workers(num_workers), @@ -1232,6 +1279,7 @@ def set_batch_size(self, new_batch_size: int) -> None: batch_size=batch_size, shuffle=shuffle, num_workers=num_workers, + worker_init_fn=_close_inherited_h5_fds, drop_last=drop_last, pin_memory=pin_memory, collate_fn=collate_fn, diff --git a/weightslab/backend/logger.py b/weightslab/backend/logger.py index 7ce7febe..0c2cbc3e 100644 --- a/weightslab/backend/logger.py +++ b/weightslab/backend/logger.py @@ -36,7 +36,7 @@ import os import threading import time -from collections import defaultdict +from collections import defaultdict, deque import duckdb import pandas as pd @@ -64,6 +64,18 @@ _STAGE_FLUSH_THRESHOLD = 50_000 # How often the background flush thread wakes up (see LoggerQueue._flush_loop). +def _default_history_tail() -> int: + """Recent points kept per sample for signal-DAG history reads. + + Bounded so history() costs O(batch) instead of scanning the whole + per_sample table (140ms at 20M rows, once per step, and growing). + """ + try: + return max(0, int(os.environ.get("WL_HISTORY_TAIL", "16"))) + except (TypeError, ValueError): + return 16 + + def _default_flush_interval_seconds() -> float: try: return float(os.environ.get("WL_LOGGER_FLUSH_INTERVAL_SECONDS", "2.0")) @@ -276,6 +288,10 @@ def __init__(self, register: bool = True, db_path: str = ":memory:") -> None: _qps_maxsize = int(os.environ.get("WL_QUERY_CACHE_MAXSIZE", "2048")) self._qps_version: dict = defaultdict(int) self._qps_cache_step: int = -1 + # {signal: {sample_id: deque}} -- the recent tail of each sample's + # per-sample values, maintained on write so history() never scans. + self._tail_len = _default_history_tail() + self._recent_tail: dict = defaultdict(dict) self._qps_cache = functools.lru_cache(maxsize=_qps_maxsize)(self._query_per_sample_uncached) self._qps_step_cache = functools.lru_cache(maxsize=_qps_maxsize)(self._query_per_sample_at_step_uncached) @@ -757,6 +773,20 @@ def _next_seq(self) -> int: self._seq += 1 return s + def recent_per_sample(self, graph_name: str, sample_ids): + """Recent in-memory values per sample: ``{sample_id: [values]}``. + + O(batch). Samples not written in this process are absent rather than + empty-listed; callers treat both as "not enough history yet". + """ + tail = self._recent_tail.get(graph_name) or {} + out = {} + for s in sample_ids: + q = tail.get(str(s)) + if q: + out[s] = list(q) + return out + def _maybe_autoflush(self) -> None: if (len(self._stage_signals) + len(self._stage_sample) + len(self._stage_instance)) >= _STAGE_FLUSH_THRESHOLD: @@ -823,6 +853,13 @@ def _stage_sample_row(self, graph_name, exp_hash, sample_id, step, value): ) self._qps_version[graph_name] += 1 # invalidate this signal's cached reads self._loss_shape_dirty_samples[(graph_name, exp_hash)].add(str(sample_id)) + if self._tail_len: + _tail = self._recent_tail[graph_name] + _sid = str(sample_id) + _q = _tail.get(_sid) + if _q is None: + _q = _tail[_sid] = deque(maxlen=self._tail_len) + _q.append(float(value)) self._maybe_autoflush() def _invalidate_qps_cache(self) -> None: diff --git a/weightslab/cli.py b/weightslab/cli.py index d4e5adb4..65a592e1 100644 --- a/weightslab/cli.py +++ b/weightslab/cli.py @@ -623,6 +623,31 @@ def example_start(args): try: env = os.environ.copy() env['WEIGHTSLAB_SUPPRESS_BANNER'] = '1' + # `weightslab start` runs in its own terminal, so its + # WEIGHTSLAB_ROOT_LOG_DIR export never reaches this process -- read the + # directory it recorded and hand it to the example, so the run lands in + # the experiment the UI is showing instead of a throwaway temp dir. + # Anything already set in this shell wins: an explicit choice by the + # user must not be overridden by the last UI launch. + if not (env.get('WEIGHTSLAB_ROOT_LOG_DIR') or '').strip(): + try: + from weightslab.utils.active_experiment import live_ui_experiment_dir + adopted = live_ui_experiment_dir() + except Exception as exc: # noqa: BLE001 + logger.debug(f"Could not read the active experiment directory: {exc}") + adopted = None + if adopted: + env['WEIGHTSLAB_ROOT_LOG_DIR'] = adopted + logger.info(f" Using the experiment directory from `weightslab start`: {adopted}") + else: + logger.warning( + " No experiment directory found (no WEIGHTSLAB_ROOT_LOG_DIR here and no " + "`weightslab start` on record) — this example will write to a temporary " + "directory, and the UI will not find its reports or notebook. Start the UI " + "first with `weightslab start`, or set WEIGHTSLAB_ROOT_LOG_DIR." + ) + else: + logger.info(f" Using WEIGHTSLAB_ROOT_LOG_DIR from this shell: {env['WEIGHTSLAB_ROOT_LOG_DIR']}") result = subprocess.run([sys.executable, str(main_py)], cwd=str(example_dir), env=env) except KeyboardInterrupt: logger.info("Example stopped.") @@ -892,6 +917,18 @@ def ui_start_native(args): experiment_dir = _resolve_experiment_dir(getattr(args, "experiment_dir", None)) os.environ["WEIGHTSLAB_ROOT_LOG_DIR"] = str(experiment_dir) os.environ["WL_LAST_EXPERIMENT_DIR"] = str(experiment_dir) + # An environment variable reaches only THIS process and its children. A + # training run started from another terminal is a different process tree, + # so also record the directory in the marker file every later weightslab + # process reads (weightslab.utils.active_experiment) -- without it, such a + # run fell through to a throwaway %TEMP% directory while this UI listed an + # empty reports/ from the directory established here. + try: + from weightslab.utils.active_experiment import record_ui_experiment + record_ui_experiment(experiment_dir) + except Exception as exc: # noqa: BLE001 -- advisory record, never fatal + logger.debug(f"Could not record the active experiment directory: {exc}") + # Re-recorded below with the resolved ports, once they are known. _print_experiment_guidance(experiment_dir) # If the agent has been initialized, provision OpenCode up front (in the @@ -955,6 +992,15 @@ def ui_start_native(args): ui_port = actual_port logger.info(f"UI port source: {ui_port_source} (preferred {preferred_ui_port}, using {ui_port})") os.environ["WL_LAST_UI_PORT"] = str(ui_port) + # Now that the ports are settled, stamp them on this UI's record: the + # backend port is what lets this UI ask for the experiment directory of + # ITS backend rather than of whichever backend started last -- the + # difference that matters when two experiments run side by side. + try: + from weightslab.utils.active_experiment import record_ui_experiment + record_ui_experiment(experiment_dir, ui_port=ui_port, backend_port=backend_port) + except Exception as exc: # noqa: BLE001 + logger.debug(f"Could not record the UI ports: {exc}") ui_server.serve_ui( ui_host=ui_host, diff --git a/weightslab/data/dataframe_manager.py b/weightslab/data/dataframe_manager.py index 76b9a8f7..59afaf2c 100644 --- a/weightslab/data/dataframe_manager.py +++ b/weightslab/data/dataframe_manager.py @@ -92,7 +92,15 @@ def __init__(self, flush_interval: float = 3.0, flush_max_rows: int = 100, enabl self._store: H5DataFrameStore | None = None self._array_store: H5ArrayStore | None = None self._origin_revisions: Dict[str, int] = {} + # Bumped ONLY when the `discarded` column changes. The deny-list cache + # keys on this instead of _origin_revisions, which training bumps every + # step (per-sample signal writes) and which therefore never lets a cache + # hit -- turning a per-batch len() into a 3.96M-row scan. + self._discard_revisions: Dict[str, int] = {} self._pending: set[int] = set() + # Parallel dirty set for the view. _pending is drained by the H5 flush, + # so the view cannot share it without one consumer starving the other. + self._view_pending: set = set() self._force_flush = False self._flush_interval = flush_interval self._flush_max_rows = flush_max_rows @@ -301,9 +309,66 @@ def _expand_dataframe_with_annotations(self, df: pd.DataFrame) -> pd.DataFrame: if not isinstance(work.index, pd.MultiIndex) and work.index.name != SID: work = work.copy() work.index.name = SID + fast = self._expand_fast_no_instances(work) + if fast is not None: + return fast + records = work.reset_index().to_dict("records") return self._expand_records_to_multi_index(records) + def _expand_fast_no_instances(self, work: pd.DataFrame): + """Vectorized expansion for frames with no per-instance targets. + + Returns the (sample_id, annotation_id=0) frame, or None to signal the + caller must use the record-by-record path. + """ + SID = SampleStats.Ex.SAMPLE_ID.value + ANNOT = SampleStats.Ex.INSTANCE_ID.value + TARGET = SampleStats.Ex.TARGET.value + try: + flat = work.reset_index() + if SID not in flat.columns: + return None + if TARGET in flat.columns: + tgt = flat[TARGET] + # Non-object dtype cannot hold a list => every target is scalar. + if tgt.dtype == object: + for v in tgt.to_numpy(): + if isinstance(v, (list, tuple, np.ndarray)) and len(v) > 0: + return None + # _normalize_sample_id always returns str(). astype(str) reproduces + # that for numeric dtypes only -- bytes would render as "b'x'". + sid_ser = flat[SID] + if pd.api.types.is_integer_dtype(sid_ser) or pd.api.types.is_float_dtype(sid_ser): + sids = sid_ser.astype(str).tolist() + else: + sids = [self._normalize_sample_id(v) for v in sid_ser.to_numpy()] + + out = flat.drop(columns=[c for c in (SID, ANNOT) if c in flat.columns]) + # The record path goes through python lists, so extension dtypes + # (string[pyarrow], categorical) come back as object. Match it. + for c in out.columns: + if isinstance(out[c].dtype, pd.api.extensions.ExtensionDtype): + out[c] = out[c].astype(object) + out.index = pd.MultiIndex.from_arrays( + [sids, np.zeros(len(out), dtype=np.int64)], names=[SID, ANNOT]) + return out + except Exception: + return None + + def _normalize_sample_id_index(self, values) -> "pd.Index": + """Vectorized _normalize_sample_id over an Index (~2x; 0.8s -> 0.5s at 2M). + + _normalize_sample_id is str() after unwrapping numpy scalars/bytes, so + astype(str) is exact for numeric dtypes; anything else keeps the loop. + """ + try: + if pd.api.types.is_integer_dtype(values) or pd.api.types.is_float_dtype(values): + return pd.Index(values.astype(str)) + except Exception: + pass + return pd.Index([self._normalize_sample_id(v) for v in values]) + def _normalize_sample_id(self, sample_id: Any) -> Any: """Normalize incoming sample IDs while preserving numeric IDs when possible.""" try: @@ -320,6 +385,22 @@ def _normalize_sample_id(self, sample_id: Any) -> Any: return str(sample_id) + def _level0_index(self): + """Level-0 (sample_id) values of the ledger index, cached. + + Keyed on the index object's identity: pandas Index is immutable, so a + reindex or rebuild yields a new object and invalidates this. Reusing the + object also reuses its hash engine, which is what makes a membership + probe O(1) instead of O(rows). + """ + idx = self._df.index + key = id(idx) + if getattr(self, "_lvl0_key", None) != key: + self._lvl0_key = key + self._lvl0 = (idx.get_level_values(0) + if isinstance(idx, pd.MultiIndex) else idx) + return self._lvl0 + def _coerce_sample_id_for_index(self, sample_id: Any) -> Any: """Coerce sample_id to match current dataframe index representation. @@ -332,8 +413,10 @@ def _coerce_sample_id_for_index(self, sample_id: Any) -> Any: # Check if multi-index if isinstance(self._df.index, pd.MultiIndex): - # Get level 0 (sample_id level) values - level_0_values = self._df.index.get_level_values(0) + # Cached: get_level_values(0) built a new Index over every row on each + # call, and a new object means a new hash engine, so this probe was + # O(rows) per sample. + level_0_values = self._level0_index() if sid in level_0_values: return sid sid_str = str(sid) @@ -355,6 +438,11 @@ def set_array_store(self, array_store: H5ArrayStore): if self._enable_h5_persistence: self._array_store = array_store + def _bump_discard_revisions(self, origins: Sequence[Any]) -> None: + for origin in origins or []: + key = str(origin) + self._discard_revisions[key] = self._discard_revisions.get(key, 0) + 1 + def _bump_origin_revisions(self, origins: Sequence[Any]) -> None: for origin in origins: if origin is None or pd.isna(origin): @@ -720,7 +808,7 @@ def upsert_df(self, df_local: List | pd.DataFrame, origin: str = None, force_flu # Normalize sample_id values in multi-index if isinstance(df_norm.index, pd.MultiIndex) and df_norm.index.nlevels >= 1: - level_0_normalized = pd.Index([self._normalize_sample_id(v) for v in df_norm.index.get_level_values(0)]) + level_0_normalized = self._normalize_sample_id_index(df_norm.index.get_level_values(0)) try: if df_norm.index.nlevels == 2: df_norm.index = pd.MultiIndex.from_arrays( @@ -784,7 +872,29 @@ def upsert_df(self, df_local: List | pd.DataFrame, origin: str = None, force_flu for col in all_cols: if col in self._df.columns and isinstance(self._df[col].dtype, pd.CategoricalDtype): self._df[col] = self._df[col].astype(object) - self._df.loc[existing_idx, all_cols] = df_norm.loc[existing_idx, all_cols] + # Label-aligned 2D .loc realigns the entire frame (22s at 4M rows + # when adding columns to every row). Resolve row positions once, + # then write each column positionally. Falls back if the index + # has duplicates/misses, where get_indexer returns -1. + _pos = self._df.index.get_indexer(existing_idx) + if len(_pos) and (_pos >= 0).all(): + for _c in all_cols: + _ci = self._df.columns.get_loc(_c) + _vals = df_norm.loc[existing_idx, _c].to_numpy() + # An object array (None for "no value yet") written into + # a float column upcasts the WHOLE column to object and + # it never returns -- an 8x penalty on every later sort. + # Coerce to the target dtype so None becomes NaN instead. + try: + _tgt = self._df[_c].dtype + if (_vals.dtype == object + and getattr(_tgt, "kind", "") in "fiu"): + _vals = pd.to_numeric(_vals, errors="coerce") + except Exception: + pass + self._df.iloc[_pos, _ci] = _vals + else: + self._df.loc[existing_idx, all_cols] = df_norm.loc[existing_idx, all_cols] # Append rows that do not exist yet. Use a boolean mask (not # .loc[difference]) so a duplicate key in df_norm can't be @@ -800,12 +910,28 @@ def upsert_df(self, df_local: List | pd.DataFrame, origin: str = None, force_flu if df_norm[col].dtype == bool: self._df[col] = self._df[col].fillna(False).astype(bool) + # Columns SampleStatsEx documents a boolean default for must never + # hold NaN ("None are not accepted by PD H5 storage" -- see its + # DEFAULTS). Two ways they did anyway, and the loop above catches + # neither: it only visits columns being ADDED by this upsert, and + # only when the incoming slice's dtype is already bool -- which it + # is not precisely when the slice carries missing values. So a + # sample registered without the flag, and every per-annotation row + # (sample-level values live on annotation 0 only), kept a NaN. + # + # That is not cosmetic: bool(float("nan")) is True in Python, so a + # NaN flag read as a set one -- samples served to the studio as + # discarded while the dataframe said nothing was. + self._fill_documented_flag_defaults() + # Auto-register any string-valued tag: columns as categorical tags # (e.g. dataset metadata declaring tag:weather = "rainy"/"sunny"). self._auto_register_categorical_tags(df_norm) # Optimize memory by converting repetitive columns to categorical - self._df = self._optimize_dataframe_memory(self._df) + # Only columns just written can have changed dtype-wise; a full-frame + # nunique() over every object column was ~14s of a 670s startup. + self._df = self._optimize_dataframe_memory(self._df, columns=set(df_norm.columns)) # Mark dirty for flush (handle multi-index) if isinstance(df_norm.index, pd.MultiIndex): @@ -814,6 +940,8 @@ def upsert_df(self, df_local: List | pd.DataFrame, origin: str = None, force_flu sample_ids = df_norm.index.tolist() self.mark_dirty_batch(sample_ids, force_flush=force_flush) self._bump_origin_revisions(affected_origins) + if SampleStats.Ex.DISCARDED.value in set(df_norm.columns): + self._bump_discard_revisions(affected_origins) def mark_dirty(self, sample_id: int): """Mark sample as dirty for H5 flush. @@ -823,6 +951,7 @@ def mark_dirty(self, sample_id: int): with self._lock: normalized_id = self._coerce_sample_id_for_index(sample_id) self._pending.add(normalized_id) + self._view_pending.add(normalized_id) def drop_column(self, column: str): with self._lock: @@ -833,6 +962,7 @@ def drop_column(self, column: str): def mark_dirty_batch(self, sample_ids: List[int], force_flush: bool = False): with self._lock: self._pending.update(set(sample_ids)) + self._view_pending.update(set(sample_ids)) if force_flush: self._force_flush = True @@ -1337,9 +1467,55 @@ def update_values(self, origin: str, sample_id: int, updates: Dict[str, Any], an self._df = pd.concat([self._df, df_local]) self._bump_origin_revisions([origin]) - def get_origin_revision(self, origin: str) -> int: + def take_view_dirty(self, limit: int | None = None): + """Drain and return the sample_ids changed since the last view sync. + + Returns None when the backlog exceeds *limit*, meaning a differential + update would cost more than a rebuild — the caller should fall back. + """ + with self._lock: + if limit is not None and len(self._view_pending) > limit: + # Do NOT clear here. The caller is expected to rebuild, but that + # rebuild can bail (contended update lock) -- and these ids would + # then be lost with nothing to re-mark them. clear_view_dirty() + # is called once the rebuilt view is actually swapped in. + return None + out = list(self._view_pending) + self._view_pending.clear() + return out + + def clear_view_dirty(self): + """Drop the view-dirty backlog: a full rebuild has made the view current.""" + with self._lock: + self._view_pending.clear() + + def get_source_rows(self, sample_ids, columns=None): + """Rows for *sample_ids* straight from the source frame. O(len(ids)).""" with self._lock: - return int(self._origin_revisions.get(str(origin), 0)) + if self._df.empty or not len(sample_ids): + return None + idx = self._df.index + keys = idx.get_level_values(0) if isinstance(idx, pd.MultiIndex) else idx + want = set(str(s) for s in sample_ids) + mask = keys.astype(str).isin(want) + sub = self._df.loc[mask] + return sub[columns] if columns else sub + + def get_origin_revision(self, origin: str) -> int: + # No lock: a dict read is atomic under the GIL, and this is polled from + # __len__ on every batch. Taking self._lock here put the training thread + # in contention with the flush thread on the hot path. + return int(self._origin_revisions.get(str(origin), 0)) + + def get_discard_revision(self, origin: str) -> int: + """Revision of the `discarded` column for *origin*. + + Changes only on discard/restore, so a consumer that depends purely on + deny-list state can cache against it across training steps. Read without + the lock: called per batch, and a stale-by-one read is harmless (the next + batch picks the change up), whereas lock contention here is not. + """ + return int(self._discard_revisions.get(str(origin), 0)) def update_by_groups_bulk(self, origin: str, group_ids: List[Any], updates_list: List[Dict[str, Any]]): """Broadcast updates to multiple groups in one pass.""" @@ -1535,22 +1711,42 @@ def get_sample_column_values(self, sample_ids: List[Any], column: str) -> Dict[A coerced_ids = [self._coerce_sample_id_for_index(sid) for sid in sample_ids] + idx = self._df.index try: - if isinstance(self._df.index, pd.MultiIndex) and self._df.index.nlevels >= 2: - sample_level = self._df.index.get_level_values(0) - anno_level = self._df.index.get_level_values(1) - mask = sample_level.isin(coerced_ids) & (anno_level == 0) - slice_df = self._df[mask] - sids = slice_df.index.get_level_values(0) + # _positional NB_SEEN lookup: the wanted rows are exactly + # (sample_id, 0), so resolve their POSITIONS instead of scanning. + # The masked form below materialised both index levels and copied + # a boolean-masked frame over every row -- ~1.2s/step at 3.96M + # just to read a batch of integers. + if idx.has_duplicates: + raise ValueError("non-unique index; use the scan path") + if isinstance(idx, pd.MultiIndex) and idx.nlevels >= 2: + pos = idx.get_indexer([(cid, 0) for cid in coerced_ids]) else: - mask = self._df.index.isin(coerced_ids) - slice_df = self._df[mask] - sids = slice_df.index - - for sid, val in zip(sids, slice_df[column]): - values[self._normalize_sample_id(sid)] = val + pos = idx.get_indexer(list(coerced_ids)) + col = self._df[column].to_numpy() + for cid, p in zip(coerced_ids, pos): + if p >= 0: + values[self._normalize_sample_id(cid)] = col[p] except Exception: - pass + # Fallback: original scan, for a non-unique index or any dtype + # mismatch get_indexer will not tolerate. + try: + if isinstance(idx, pd.MultiIndex) and idx.nlevels >= 2: + sample_level = idx.get_level_values(0) + anno_level = idx.get_level_values(1) + mask = sample_level.isin(coerced_ids) & (anno_level == 0) + slice_df = self._df[mask] + sids = slice_df.index.get_level_values(0) + else: + mask = idx.isin(coerced_ids) + slice_df = self._df[mask] + sids = slice_df.index + + for sid, val in zip(sids, slice_df[column]): + values[self._normalize_sample_id(sid)] = val + except Exception: + pass return values @@ -1907,7 +2103,9 @@ def _apply_buffer_records(self, records: List[Dict[str, Any]]): self._apply_updates_frame_locked(instance_df, broadcast=False) self._apply_updates_frame_locked(sample_df, broadcast=True) # Keep newly-added signal columns float32 and empty object cells as None. - self._df = self._optimize_dataframe_memory(self._df) + self._df = self._optimize_dataframe_memory( + self._df, + columns=set(sample_df.columns) | set(instance_df.columns)) # Mark all as pending for h5 flush (outside lock) self.mark_dirty_batch(sample_ids) @@ -1945,13 +2143,22 @@ def _apply_buffer_records_nonblocking(self, records: List[Dict[str, Any]]): applied_index = written_s.append(written_i) if len(written_i) else written_s update_cols = sample_df.columns.union(instance_df.columns) # Keep newly-added signal columns float32 and empty object cells as None. - _df = self._optimize_dataframe_memory(self._df) + _df = self._optimize_dataframe_memory(self._df, columns=set(update_cols)) self._df = _df finally: self._lock.release() # Det→seg conversion / array normalization over all written rows. - if applied_index is not None and len(applied_index) > 0: + # This prepares cells for the H5 write, so it is only worth doing for + # columns that write will actually include. When predictions/targets are + # excluded from the save list (WEIGHTSLAB_SAVE_PREDICTIONS_IN_H5=0) the + # pass would otherwise call get_mask per row -- which reads the source + # image to size the mask -- for arrays that are never persisted. + _savable = set(_filter_columns_by_patterns( + list(update_cols), SAMPLES_STATS_TO_SAVE_TO_H5)) + _norm_cols = [c for c in update_cols + if c in self._array_columns and c in _savable] + if applied_index is not None and len(applied_index) > 0 and _norm_cols: if applied_index.has_duplicates: applied_index = applied_index[~applied_index.duplicated()] normalized_rows = self._df.loc[applied_index].apply( @@ -2008,6 +2215,33 @@ def _flush_to_h5_if_needed(self, force: bool = False, blocking: bool = False): # Everything below happens WITHOUT locks - fully async self._flush_snapshot_to_h5(data_snapshot, work) + def _rows_with_array_cells(self, data_snapshot: pd.DataFrame): + """Index labels of rows that may hold an array-valued cell. + + Column-wise scan replacing a full iterrows() pass: only object-dtype + columns can hold an ndarray/list/tuple/ArrayH5Proxy, and rows with none + of those are no-ops for the caller. + """ + cols = [c for c in self._array_columns if c in data_snapshot.columns] + if not cols: + return [] + hits = None + for col in cols: + ser = data_snapshot[col] + if ser.dtype != object: + continue # a numeric column cannot hold an array object + vals = ser.to_numpy() + mask = np.fromiter( + (isinstance(v, (np.ndarray, list, tuple, ArrayH5Proxy)) for v in vals), + dtype=bool, count=len(vals)) + if not mask.any(): + continue + found = ser.index[mask] + hits = found if hits is None else hits.union(found) + if hits is None: + return [] + return list(hits) + def _flush_snapshot_to_h5(self, data_snapshot: pd.DataFrame, work: List[int]): """Flush data snapshot to H5 - runs completely outside locks. @@ -2030,7 +2264,10 @@ def _flush_snapshot_to_h5(self, data_snapshot: pd.DataFrame, work: List[int]): arrays_to_store: Dict[str, Dict[str, np.ndarray]] = {} rowkey_to_index: Dict[str, Any] = {} - for idx, row in data_snapshot.iterrows(): + for idx in self._rows_with_array_cells(data_snapshot): + row = data_snapshot.loc[idx] + if isinstance(row, pd.DataFrame): # duplicate label guard + row = row.iloc[0] if is_multi: sample_id, annot = idx[0], int(idx[1]) else: @@ -2103,7 +2340,7 @@ def _flush_snapshot_to_h5(self, data_snapshot: pd.DataFrame, work: List[int]): except Exception as e: logger.error(f"[LedgeredDataFrameManager] Error flushing to H5: {e}") - def _optimize_dataframe_memory(self, df: pd.DataFrame, categorical_tags: Dict[str, List[str]] | None = None) -> pd.DataFrame: + def _optimize_dataframe_memory(self, df: pd.DataFrame, categorical_tags: Dict[str, List[str]] | None = None, columns=None) -> pd.DataFrame: """Optimize dataframe memory by converting repetitive string columns to categorical. Categorical dtype compresses repeated values: instead of storing each string, @@ -2130,12 +2367,39 @@ def _optimize_dataframe_memory(self, df: pd.DataFrame, categorical_tags: Dict[st if categorical_tags is None: categorical_tags = self._categorical_tags + # Honour `columns`: only the columns this flush actually wrote can have + # gained a float64 signal value or a fresh NaN, so scanning the rest is + # O(rows) of pure waste on every flush. + _scan_cols = (list(df.columns) if columns is None + else [c for c in df.columns if c in columns]) + + # === 0) Repair signal columns that were upcast to object === + # A single None written into a float column converts it permanently, and + # object columns sort ~8x slower. signals//* are numeric by definition, + # so any object one is damage rather than intent. Runs BEFORE the + # NaN->None pass below, which only touches object columns and would + # otherwise keep them that way. + for col in _scan_cols: + if not str(col).startswith("signals") or df[col].dtype != object: + continue + try: + coerced = pd.to_numeric(df[col], errors="coerce") + # Only if nothing was lost: a genuine non-numeric value means the + # column is not what we think it is, so leave it alone. + if coerced.notna().sum() == df[col].notna().sum(): + df[col] = coerced.astype(np.float32) + logger.debug( + "[LedgeredDataFrameManager] restored '%s' object -> float32", col) + except Exception as exc: + logger.debug("[LedgeredDataFrameManager] dtype repair skipped for '%s': %s", + col, exc) + # === 1) Downcast float64 signal columns to float32 === # Per-sample / per-instance signal scalars (loss & metric values) don't need # float64 precision, so this halves the cost of every ``signals//*`` column # with no practical loss for monitoring. Numeric dtype is preserved (NaNs # stay NaN). Done before categorical conversion below. - for col in df.columns: + for col in _scan_cols: if str(col).startswith("signals") and df[col].dtype == np.float64: try: df[col] = df[col].astype(np.float32) @@ -2159,7 +2423,12 @@ def _optimize_dataframe_memory(self, df: pd.DataFrame, categorical_tags: Dict[st # into a Categorical, its missing cells are stuck as NaN (categorical code # -1) and can no longer be set to a plain None — so we clean object cells to # None here first, while they are still plain object dtype. - for col in df.columns: + for col in _scan_cols: + # The docstring above says numeric/bool/categorical are skipped, but the + # loop had no dtype check: at registration every signals//* column is all + # NaN, so this did a full label-aligned write per float column. + if df[col].dtype != object: + continue na_mask = df[col].isna() if na_mask.any(): df.loc[na_mask, col] = None @@ -2171,6 +2440,17 @@ def _optimize_dataframe_memory(self, df: pd.DataFrame, categorical_tags: Dict[st SampleStats.Ex.TASK_TYPE.value, # Task type (e.g. 'classification', 'segmentation') ] + # nunique() over a 4M-row MultiIndex costs ~8s AND holds the GIL, stalling + # the training thread inside forward/backward. Only an object-dtype + # candidate reads it, so compute it on first use rather than every flush. + _n_rows_cache = [] + + def _n_rows(): + if not _n_rows_cache: + _n_rows_cache.append( + df.index.get_level_values(0).nunique() + if isinstance(df.index, pd.MultiIndex) else len(df)) + return _n_rows_cache[0] for col in categorical_candidates: if col not in df.columns: continue @@ -2182,15 +2462,14 @@ def _optimize_dataframe_memory(self, df: pd.DataFrame, categorical_tags: Dict[st # Only convert to categorical if: # 1. Column contains strings (object dtype) # 2. Number of unique values < 50% of total rows (good compression ratio) + if columns is not None and col not in columns: + continue if df[col].dtype == 'object': n_unique = df[col].nunique() # Use unique sample count, not row count — with MultiIndex each # sample has multiple annotation rows which would inflate n_rows # and make the ratio appear better than it really is. - if isinstance(df.index, pd.MultiIndex): - n_rows = df.index.get_level_values(0).nunique() - else: - n_rows = len(df) + n_rows = _n_rows() compression_ratio = n_unique / n_rows if n_rows > 0 else 1.0 if compression_ratio < 0.5 and n_unique > 1: # Worth compressing if < 50% unique @@ -2319,6 +2598,44 @@ def get_combined_df( return df + def _fill_documented_flag_defaults(self) -> None: + """Give every boolean column SampleStatsEx documents a default its default. + + Only the columns with a ``bool`` default in + ``SampleStatsEx.DEFAULTS`` (today: ``discarded``) -- a ``tag:*`` column + is deliberately left sparse, where NaN and False mean the same thing + and NaN costs nothing to store, and a *categorical* tag's NaN means + "unset", which is not a default at all. + + Cheap by design: an ``isna().any()`` short-circuit per flag column, so + the common case (nothing missing) touches no rows. + """ + if self._df is None or self._df.empty: + return + # DEFAULTS lives on SampleStats, the outer class -- SampleStatsEx is + # only its `Ex` enum, so reading it off that is a silent no-op. + from weightslab.data.sample_stats import SampleStats + defaults = getattr(SampleStats, "DEFAULTS", {}) or {} + for col, default in defaults.items(): + if not isinstance(default, bool) or col not in self._df.columns: + continue + try: + series = self._df[col] + if not series.isna().any(): + continue + if isinstance(series.dtype, pd.CategoricalDtype): + # fillna on a Categorical raises unless the value is a + # known category; widen first, then let the memory pass + # re-apply the dtype. + series = series.astype(object) + self._df[col] = series.fillna(default) + logger.debug( + "[LedgeredDataFrameManager] filled missing %r with its " + "documented default %r", col, default) + except Exception as exc: # noqa: BLE001 -- never fail an upsert on this + logger.debug( + "[LedgeredDataFrameManager] could not default %r: %s", col, exc) + def get_collapse_annotations_to_samples_df(self, df: pd.DataFrame | None = None) -> pd.DataFrame: """Collapse a (sample_id, annotation_id) multi-index df to one row per sample. diff --git a/weightslab/data/h5_array_store.py b/weightslab/data/h5_array_store.py index 02c3e7e3..405cf02b 100644 --- a/weightslab/data/h5_array_store.py +++ b/weightslab/data/h5_array_store.py @@ -514,6 +514,52 @@ def save_array( finally: self._rw_lock.release_write() + def _try_inplace_batch(self, prepared): + """Overwrite existing datasets in place; None means "cannot, fall back". + + Two passes under the write lock: check every destination exists with a + matching shape and dtype, and only then write. A partial in-place write + followed by a fallback would corrupt silently, so nothing is written + until the whole batch is known to fit. + """ + if not self._path.exists(): + return None + with self._local_lock: + self._rw_lock.acquire_write() + try: + with _InterProcessFileLock(self._lock_path, timeout=self._lock_timeout, + poll_interval=self._poll_interval): + with h5py.File(str(self._path), 'a') as f: + for group_name, key_data in prepared.items(): + grp = f.get(group_name) + if grp is None: + return None + for key_name, (array, _meta) in key_data.items(): + kg = grp.get(key_name) + if kg is None or 'data' not in kg: + return None + dset = kg['data'] + if dset.shape != array.shape or dset.dtype != array.dtype: + return None + for group_name, key_data in prepared.items(): + for key_name, (array, metadata) in key_data.items(): + kg = f[group_name][key_name] + kg['data'][...] = array + for mk, mv in metadata.items(): + kg.attrs[mk] = mv + return { + group_name: { + key_name: self._build_path_reference(group_name, key_name) + for key_name in key_data + } + for group_name, key_data in prepared.items() + } + except Exception as exc: + logger.debug(f"[H5ArrayStore] in-place batch fell back: {exc}") + return None + finally: + self._rw_lock.release_write() + def save_arrays_batch( self, arrays_dict: Dict[int, Dict[str, np.ndarray]], @@ -562,6 +608,15 @@ def save_arrays_batch( if not prepared: return {} + # O(change): if every array already exists with the same shape and dtype, + # overwrite the values in place. That is not a structural change, so it + # needs neither the temp file nor the full-file backup (9.5GB per flush + # at current ledger size). Returns None if anything would need creating + # or resizing, and the original two-phase path below runs unchanged. + inplace_refs = self._try_inplace_batch(prepared) + if inplace_refs is not None: + return inplace_refs + tmp_path = self._path.with_suffix(f".h5.writing_{uuid.uuid4().hex[:8]}") try: with h5py.File(str(tmp_path), 'w') as f_tmp: diff --git a/weightslab/data/h5_dataframe_store.py b/weightslab/data/h5_dataframe_store.py index 476b5655..431425d1 100644 --- a/weightslab/data/h5_dataframe_store.py +++ b/weightslab/data/h5_dataframe_store.py @@ -680,6 +680,120 @@ def load_all(self, origins: Iterable[str] = None, columns: Optional[Iterable[str return pd.DataFrame() raise + def ensure_index(self, origin: str, columns=("sample_id",)) -> bool: + """Build the on-disk column index deliberately (checkpoint / first query). + + Kept OFF the flush path: rebuilding it per upsert costs 92.7s at 4M rows + versus 6.9s without, for an index no hot-path read uses. + """ + key = self._key(origin) + try: + with self._local_lock: + with _InterProcessFileLock(self._lock_path, timeout=self._lock_timeout, + poll_interval=self._poll_interval): + with pd.HDFStore(str(self._path), mode="a") as store: + if key not in store: + return False + store.create_table_index(key, columns=list(columns), + optlevel=6, kind="medium") + return True + except Exception as exc: + logger.warning(f"[H5DataFrameStore] ensure_index({origin}) failed: {exc}") + return False + + # --- O(change) in-place update ------------------------------------------ + _POSMAP_CACHE: dict = {} + + def _posmap(self, store, key, force=False): + """(sample_id, annotation_id) -> row position. Rows are registered once + and never deleted, so positions are stable; built from ONE column (2.9s + at 4M rows) rather than reading the table.""" + ck = (str(self._path), key) + if not force and ck in self._POSMAP_CACHE: + return self._POSMAP_CACHE[ck] + try: + sids = store.select_column(key, "sample_id").values + try: + aids = store.select_column(key, "annotation_id").values + except Exception: + aids = np.zeros(len(sids), dtype="i8") + # Hoist the normalisation out of the insert loop: the per-row + # decode/str/int calls interleaved with dict inserts cost 3.6s at 4M + # rows, against ~1.7s for dict(zip(...)) over pre-normalised lists. + # One pass that is also the type check: str hits the identity branch, + # so a mixed-dtype column stays correct without a second full scan. + sid_list = [s if type(s) is str + else (s.decode() if isinstance(s, bytes) else str(s)) + for s in sids.tolist()] + m = dict(zip(zip(sid_list, aids.tolist()), range(len(sid_list)))) + self._POSMAP_CACHE[ck] = m + return m + except Exception as exc: + logger.debug(f"[H5DataFrameStore] posmap build failed: {exc}") + return None + + def _invalidate_posmap(self, key): + self._POSMAP_CACHE.pop((str(self._path), key), None) + + def _try_inplace(self, store, key, df_norm) -> bool: + """Overwrite existing rows' values in place. True if fully applied.""" + try: + import tables as _tables + except Exception: + return False + try: + node = store._handle.get_node(key) + tbl = getattr(node, "table", node) + if not isinstance(tbl, _tables.Table): + return False + # An indexed column cannot be modified in place (PyTables raises). + if any(tbl.cols._f_col(c).is_indexed for c in tbl.colnames): + return False + + cols = [c for c in df_norm.columns if c in tbl.colnames] + if len(cols) != len(df_norm.columns): + return False # new column => schema change + + pos = self._posmap(store, key) + if not pos: + return False + + idx = df_norm.index + if isinstance(idx, pd.MultiIndex): + pairs = [(str(a), int(b)) for a, b in zip(idx.get_level_values(0), + idx.get_level_values(1))] + else: + pairs = [(str(a), 0) for a in idx] + + coords = np.empty(len(pairs), dtype=np.int64) + for i, p in enumerate(pairs): + j = pos.get(p) + if j is None: + return False # unknown row => not an update + coords[i] = j + + order = np.argsort(coords) # PyTables wants ascending coords + coords_sorted = coords[order] + rec = tbl.read_coordinates(coords_sorted) + for c in cols: + vals = df_norm[c].to_numpy()[order] + tgt = rec[c].dtype + if tgt.kind == "S": + vals = np.array([("" if v is None else str(v)).encode()[:tgt.itemsize] + for v in vals], dtype=tgt) + else: + try: + vals = vals.astype(tgt, copy=False) + except Exception: + return False # dtype mismatch => fall back + rec[c] = vals + tbl.modify_coordinates(coords_sorted, rec) + tbl.flush() + return True + except Exception as exc: + logger.debug(f"[H5DataFrameStore] in-place update fell back: {exc}") + return False + def upsert(self, origin: str, df: pd.DataFrame) -> int: """Atomic upsert with corruption prevention via backup and checksum verification.""" df_norm = self._normalize_for_write(df) @@ -689,13 +803,26 @@ def upsert(self, origin: str, df: pd.DataFrame) -> int: key = self._key(origin) self._ensure_parent() - # Create backup BEFORE any writes - backup_path = self._create_backup() + # The backup is a full copy of the file (696MB at 4M rows). Only the + # read-merge-rewrite path below can destroy the table -- the in-place + # path just overwrites values in already-allocated rows -- so the copy + # is deferred until we know we are taking the destructive route. + backup_path = None with self._local_lock: with _InterProcessFileLock(self._lock_path, timeout=self._lock_timeout, poll_interval=self._poll_interval): try: with pd.HDFStore(str(self._path), mode="a") as store: + # O(change): if every row already exists and the schema is + # unchanged, overwrite values in place (0.1ms vs 42s). + if key in store and self._try_inplace(store, key, df_norm): + return len(df_norm) + + # Nothing has been written yet in this call; flush so the + # copy below captures a consistent on-disk file. + store.flush() + backup_path = self._create_backup() + existing = pd.DataFrame() # Try to load existing data. A ValueError can surface from a @@ -801,7 +928,8 @@ def upsert(self, origin: str, df: pd.DataFrame) -> int: store.remove(key) # Write new data - store.append(key, existing, format="table", data_columns=True) + store.append(key, existing, format="table", data_columns=True, index=False) + self._invalidate_posmap(key) # Force flush to disk store.flush() @@ -880,7 +1008,7 @@ def delete_column(self, column_name: str, origins: Optional[Iterable[str]] = Non # Remove old key and write updated dataframe store.remove(key) if not df.empty: - store.append(key, df, format="table", data_columns=True) + store.append(key, df, format="table", data_columns=True, index=False) modified_count += 1 logger.debug(f"[H5DataFrameStore] Deleted column {column_name} from {origin}") diff --git a/weightslab/examples/PyTorch/wl-video-generation/utils/data.py b/weightslab/examples/PyTorch/wl-video-generation/utils/data.py index ddb4d3ab..42112fba 100644 --- a/weightslab/examples/PyTorch/wl-video-generation/utils/data.py +++ b/weightslab/examples/PyTorch/wl-video-generation/utils/data.py @@ -31,7 +31,6 @@ """ import csv import logging -import os import subprocess import shutil import wave diff --git a/weightslab/src.py b/weightslab/src.py index 88067e30..bf3c3e3e 100644 --- a/weightslab/src.py +++ b/weightslab/src.py @@ -83,7 +83,16 @@ def _resolve_configured_root_log_dir(configured): actually points at an existing directory; if it's set but stale/typo'd, a warning is logged and resolution falls through to (3) instead of silently training into a directory the UI never established. - 3. A throwaway ``tempfile.mkdtemp()`` — last resort so serving never fails + 3. The directory a RUNNING ``weightslab start`` established, read from + the marker file (see weightslab.utils.active_experiment). Only while + that UI is alive: a record left by one that has since exited must not + redirect an unrelated run. The + environment variable only reaches processes started FROM that same + shell; a training run launched in another terminal (or by + ``weightslab start example``) is a different process tree and used to + fall straight through to (4), landing in %TEMP% while the UI listed an + empty reports/ from the directory it had established. + 4. A throwaway ``tempfile.mkdtemp()`` — last resort so serving never fails for lack of a directory. """ if configured: @@ -94,9 +103,28 @@ def _resolve_configured_root_log_dir(configured): return env_dir logger.warning( f"WEIGHTSLAB_ROOT_LOG_DIR is set to '{env_dir}', but that directory " - "does not exist. Falling back to a temporary directory instead." + "does not exist. Falling back to the experiment directory recorded " + "by `weightslab start`, then to a temporary directory." ) - return tempfile.mkdtemp() + try: + from weightslab.utils.active_experiment import live_ui_experiment_dir + marker_dir = live_ui_experiment_dir() + except Exception: # noqa: BLE001 -- never block serving on the marker + marker_dir = None + if marker_dir: + logger.info( + "Using the experiment directory established by `weightslab start`: " + "%s (no root_log_dir configured and WEIGHTSLAB_ROOT_LOG_DIR is not " + "set in this process).", marker_dir) + return marker_dir + tmp_dir = tempfile.mkdtemp() + logger.warning( + "No root_log_dir configured, no WEIGHTSLAB_ROOT_LOG_DIR, and no " + "experiment directory recorded by `weightslab start` — this run will " + "write to the throwaway directory %s. Checkpoints, reports and the " + "notebook will NOT be where the UI looks for them. Start the UI first " + "(`weightslab start`), or set root_log_dir in your config.", tmp_dir) + return tmp_dir # Get global dataframe proxy (auto-updated when ledger registers real manager) @@ -494,6 +522,16 @@ def history(self, signal_name): out = {s: [] for s in self.sample_ids} if self.logger is None: return out + # Read the bounded in-memory tail: O(batch). The full-history query this + # replaced scanned the entire per_sample table on every call (140ms at + # 20M rows, once per step, growing without bound). + if not os.environ.get("WL_HISTORY_FROM_DB"): + recent = self.logger.recent_per_sample(signal_name, self.sample_ids) + for s in self.sample_ids: + vals = recent.get(s) + if vals: + out[s] = list(vals) + return out # query_per_sample accepts a list of ids -> one scan for the whole batch. for sid, step, val, _ in self.logger.query_per_sample(signal_name, sample_ids=self.sample_ids): out.setdefault(int(sid), []).append(val) # rows already ordered by seq (= step order) @@ -1023,9 +1061,23 @@ def wrappered_fwd(original_forward, kwargs, reg_name, *a, **kw): # batch and call the signal once. It returns a length-B # array. Avoids B Python calls + B SignalContext allocs, # and lets the signal do batched ledger reads. + # inputs= must be populated here too: on the + # subscribe_to path the context was built without it, + # so b.inputs was {} and any signal declaring + # inputs=[...] raised KeyError on every call. The + # subscribed signal IS the declared input here, and + # its per-sample values are already in val_vec. + _decl = meta.get('inputs') or [] + _sub = meta.get('subscribe_to') + _vals = [float(v) for v in val_vec] + _bin = {} + for _d in _decl: + if _sub is None or _d == _sub or _d == reg_name: + _bin[_d] = _vals bctx = BatchSignalContext( sample_ids=[int(u) for u in ids_np], - subscribed_values=[float(v) for v in val_vec], + subscribed_values=_vals, + inputs=_bin, logger=_lg, dataframe=df_proxy, origin=kwargs.get('origin', 'train'), @@ -1446,6 +1498,14 @@ def new_forward(*a, **kw): # _resolve_configured_root_log_dir for the resolution order). _hp_cfg['root_log_dir'] = _resolve_configured_root_log_dir( _hp_cfg.get('root_log_dir')) + # Publish where training ACTUALLY writes, so the UI lists this + # run's reports/notebooks even when the two resolved their + # directory by different routes. + try: + from weightslab.utils.active_experiment import record_backend_experiment + record_backend_experiment(_hp_cfg['root_log_dir']) + except Exception as _exc: # noqa: BLE001 -- advisory only + logger.debug("Could not record the backend experiment dir: %s", _exc) try: # Check if a checkpoint manager is already registered in ledger try: @@ -1738,6 +1798,25 @@ def serve(serving_cli: bool = True, serving_grpc: bool = True, _notebook_service.configure_embedded_kernel(embed_kernel_decision) if serving_grpc: + # Stamp the gRPC port on this backend's active-experiment record, so a + # UI proxying to that port finds THIS experiment's directory rather + # than whichever backend happened to start last (see + # active_experiment._sole_live_dir). Advisory: never fail serving on it. + try: + from weightslab.utils.active_experiment import record_backend_experiment + _grpc_port = (kwargs.get("grpc_port") + or int(os.getenv("GRPC_BACKEND_PORT", 50051))) + _root_log_dir = None + try: + _hp = ledgers.get_hyperparams() + _root_log_dir = _hp["root_log_dir"] if _hp is not None else None + except Exception: # noqa: BLE001 -- no hyperparameters registered + _root_log_dir = None + if _root_log_dir: + record_backend_experiment(_root_log_dir, grpc_port=int(_grpc_port)) + except Exception as _exc: # noqa: BLE001 + logger.debug("Could not record the backend gRPC port: %s", _exc) + grpc_serve(**kwargs) if embed_kernel_decision: @@ -4935,6 +5014,16 @@ def resolve_signal_classifier(signal_name): return _GLOBAL_CLASSIFIER or classify_loss_shape +_SHAPE_LABELS: dict = {} + + +def _label_counts(cache): + out = {} + for lab in cache.values(): + out[lab] = out.get(lab, 0) + 1 + return out + + def write_signal_shapes(signal_name, tag_name=None, classifier=None, exp_hash=None, sample_ids=None): """Reusable engine: classify each sample's own trajectory of *signal_name* into a categorical tag and return the ``{label: count}`` distribution. @@ -4953,17 +5042,29 @@ def write_signal_shapes(signal_name, tag_name=None, classifier=None, exp_hash=No clf = classifier or resolve_signal_classifier(signal_name) if tag_name is None: tag_name = signal_name + "_shape" if signal_name.endswith('_loss') else signal_name + "_loss_shape" + + # Labels for samples this pass does not reclassify are carried here, so an + # incremental call still returns a distribution over the WHOLE dataset, and + # a sample whose label is unchanged is not re-written to the ledger. + cache = _SHAPE_LABELS.setdefault(signal_name, {}) + if sample_ids is not None and not list(sample_ids): + return _label_counts(cache) + series = {} for sid, step, val, _ in query_signal_history(signal_name, exp_hash=exp_hash, sample_ids=sample_ids): series.setdefault(sid, []).append((step, val)) by_label = {} for sid, pts in series.items(): label = clf([v for _, v in sorted(pts)]) - if label is not None: - by_label.setdefault(label, []).append(sid) + if label is None: + continue + if cache.get(sid) == label: + continue # unchanged -> no ledger write needed + cache[sid] = label + by_label.setdefault(label, []).append(sid) for label, sids in by_label.items(): set_categorical_tag(sids, tag_name, label) - return {k: len(v) for k, v in by_label.items()} + return _label_counts(cache) def write_loss_shapes(loss_signal="loss_sample", classifier=None): diff --git a/weightslab/trainer/services/agent/agent.py b/weightslab/trainer/services/agent/agent.py index 7158d484..b1701b94 100644 --- a/weightslab/trainer/services/agent/agent.py +++ b/weightslab/trainer/services/agent/agent.py @@ -735,6 +735,12 @@ def _load_config(self): # actually chose is exempt from that self-healing. self._opencode_model_explicit = "OPENCODE_MODEL" in os.environ self.opencode_model = os.environ.get("OPENCODE_MODEL", "") + # agent_config.yaml's opencode_model lands here instead of pinning the + # model: it SEEDS the shared choice (used when nothing has been chosen + # yet, then published so the studio shows it), while a model picked in + # the UI afterwards wins. Pinning it meant a run started after picking + # a model in the studio silently went back to the yaml value. + self._opencode_model_seed = "" # The same directory `weightslab start ` roots the browser # landing-page agent at (WEIGHTSLAB_ROOT_LOG_DIR) -- the shared key # opencode_process.py's lock file is discovered/published under, so @@ -759,6 +765,9 @@ def _load_config(self): inner_pkg / "agent_config.yaml", Path.cwd() / "agent_config.yaml" ] + # Overwritten below only when a file is actually applied, so the banner + # can say "none found" instead of naming the last candidate it tried. + self._config_source_path = "(no agent config file found)" for path in config_paths: if not path.exists(): continue try: @@ -771,28 +780,36 @@ def _load_config(self): if a_cfg.get("opencode_url"): self._opencode_url_explicit = True self.opencode_url = a_cfg.get("opencode_url", self.opencode_url) - if a_cfg.get("opencode_model"): - self._opencode_model_explicit = True - self.opencode_model = a_cfg.get("opencode_model", self.opencode_model) + _cfg_model = str(a_cfg.get("opencode_model") or "").strip() + if _cfg_model: + # Seed, not pin -- see _opencode_model_seed above. + # OPENCODE_MODEL still wins: it is set per process, on purpose. + self._opencode_model_seed = _cfg_model + self._config_source_path = path _LOGGER.info(f"Applied agent configuration from {path}") _LOGGER.debug(f"Agent Config: {cfg}") break except Exception as e: _LOGGER.warning(f"Error loading config from {path}: {e}") - # Log the final configuration for transparency - _LOGGER.info( - "" + "\n" + - "\n# #######################################" + "\n" + - "# #######################################" + "\n" + - f"Agent initialized from configuration {path}: " + "\n" + - f"\tOpenCode URL={self.opencode_url}, Model={self.opencode_model or '(server default)'}" + "\n" + - "# #######################################" + "\n" + - "# #######################################" + "\n" + "" - ) - - def _setup_providers(self): + # The banner itself is emitted by _log_agent_configuration() AFTER + # _setup_providers has resolved the model, so it names the model + # actually in use instead of the pre-resolution blank -- which printed + # "(server default)" and read as "the studio's pick was ignored". + + def _setup_providers(self, requested_model: Optional[str] = None): + """(Re)build the OpenCode provider. + + `requested_model` is a model the USER just chose (CLI `agent model`, + `agent init --model`, the SetAgentModel RPC). It must survive this + call: the shared-config re-read below exists to follow the studio's + picker, and it used to overwrite the very model the caller had just + asked for -- `agent model X` answered "Model switched to ". A user choice is instead PUBLISHED to OpenCode's config, so it + becomes the shared choice the studio picker and every other client see + too, and is only pinned in-process when that write is refused. + """ self.chain_opencode = None self._opencode_chat = None initialized = False @@ -806,18 +823,85 @@ def _setup_providers(self): workspace_dir=self.opencode_workspace_dir, url_is_explicit=self._opencode_url_explicit, model_is_explicit=self._opencode_model_explicit, + seed_model=getattr(self, "_opencode_model_seed", ""), ) self.chain_opencode = self._opencode_chat.as_runnable() initialized = True + # Resolve up front rather than on the first query: the model is + # part of what the start-up banner reports, and a backend started + # before the studio publishes its fallback so the UI adopts the + # same model (see OpenCodeChat.resolve_model). + try: + if requested_model: + self._opencode_chat.model = requested_model + self.opencode_model = requested_model + if self._opencode_chat.publish_model(requested_model): + # Shared, not pinned: later turns keep re-reading the + # config, which now names this model -- so a studio + # pick after this still wins, as it should. + self._opencode_model_source = "user-published" + else: + # The config would not take it (read-only, older + # server). Pin it for this process so the switch the + # user asked for still takes effect. + self._opencode_chat.model_is_explicit = True + self._opencode_model_explicit = True + self._opencode_model_source = "user-pinned" + else: + resolved, source = self._opencode_chat.resolve_model(publish_default=True) + if resolved: + self.opencode_model = resolved + self._opencode_model_source = source + # _ensure_reachable may have moved us to a discovered server. + self.opencode_url = self._opencode_chat.base_url + except Exception as exc: # noqa: BLE001 -- server down; lazy path retries + _LOGGER.debug("[Agent] deferred OpenCode model resolution: %s", exc) + self._opencode_model_source = "unresolved" _LOGGER.info( f"[Agent] OpenCode enabled: {self.opencode_url} " - f"(model={self.opencode_model or 'server default'})" + f"(model={self.opencode_model or 'unresolved'})" ) except Exception as e: _LOGGER.error(f"OpenCode error: {e}") + self._log_agent_configuration() return initialized + # Human-readable provenance for the start-up banner. + _MODEL_SOURCE_LABELS = { + "pinned": "pinned by OPENCODE_MODEL", + "pinned-published": ("pinned by OPENCODE_MODEL, and published to " + "OpenCode's config so the studio shows it"), + "config-seed": ("from agent_config.yaml's opencode_model; nothing was " + "chosen in OpenCode's config yet"), + "config-seed-published": ("from agent_config.yaml's opencode_model " + "(nothing chosen yet), published to OpenCode's " + "config so the studio shows it"), + "opencode-config": "from OpenCode's config, which the studio model picker writes", + "default": "built-in default; OpenCode's config could not be updated", + "default-published": "built-in default, published to OpenCode's config for the studio", + "user-published": "chosen here and published to OpenCode's config", + "user-pinned": "chosen here; OpenCode's config refused the write, pinned to this backend", + "kept": "kept from this session", + "unresolved": "unresolved -- OpenCode unreachable, retried on the first query", + } + + def _log_agent_configuration(self) -> None: + """Start-up banner: the model actually in use, and where it came from.""" + source = getattr(self, "_opencode_model_source", "unresolved") + detail = self._MODEL_SOURCE_LABELS.get(source, source) + path = getattr(self, "_config_source_path", "(no agent config file found)") + _LOGGER.info( + "" + "\n" + + "\n# #######################################" + "\n" + + "# #######################################" + "\n" + + f"Agent initialized from configuration {path}: " + "\n" + + f"\tOpenCode URL={self.opencode_url}" + "\n" + + f"\tModel={self.opencode_model or '(unresolved)'} ({detail})" + "\n" + + "# #######################################" + "\n" + + "# #######################################" + "\n" + "" + ) + def is_available(self) -> bool: """Return True if the OpenCode provider is ready to serve requests.""" return self.chain_opencode is not None @@ -843,10 +927,12 @@ def initialize_with_cloud_key(self, api_key: str, provider: str, model: Optional if model is not None and not model.strip(): return False, "Model cannot be empty." - self.opencode_model = model.strip() if model and model.strip() else self.opencode_model + requested = model.strip() if model and model.strip() else None + if requested: + self.opencode_model = requested self.preferred_provider = "opencode" - success = self._setup_providers() + success = self._setup_providers(requested_model=requested) if self.chain_opencode is None or not success: return False, "Could not reach the OpenCode server. Please verify OPENCODE_URL and that it is running." @@ -865,11 +951,56 @@ def change_model(self, model: str) -> "tuple[bool, str]": if not model or not model.strip(): return False, "Model cannot be empty." - self.opencode_model = model.strip() - success = self._setup_providers() + requested = model.strip() + self.opencode_model = requested + success = self._setup_providers(requested_model=requested) if self.chain_opencode is None or not success: return False, "Could not reach the OpenCode server. Please verify OPENCODE_URL and that it is running." - return True, f"Model switched to {self.opencode_model}. Ready to help you." + if self.opencode_model != requested: + # Never report a switch that did not happen. + return False, (f"Could not switch to {requested}: the model in use is " + f"{self.opencode_model}.") + shared = self._opencode_model_source == "user-published" + return True, ( + f"Model switched to {self.opencode_model}. " + + ("Published to OpenCode's config, so the studio picker shows it too. " + if shared else + "OpenCode's config would not take it, so it is pinned to this backend only. ") + + "Ready to help you." + ) + + def current_model(self) -> Optional[str]: + """The model the NEXT query will actually use. + + `self.opencode_model` is only what was resolved when the provider was + last (re)initialised. A model chosen in the studio afterwards lands in + OpenCode's config, which every turn re-reads -- so reporting the + snapshot made a UI pick look ignored (`agent status` kept naming the + old model while queries already used the new one). Re-resolves here: + one local GET, on a command the user typed. + """ + chat = self._opencode_chat + if chat is None: + return self.opencode_model or None + try: + model, source = chat.resolve_model() + if model: + self.opencode_model = model + self._opencode_model_source = source + if source == "pinned": + # A pin wins for this backend, so say plainly when the studio + # is showing something else -- otherwise the two surfaces + # disagree with no explanation anywhere. + chosen = chat._configured_model() + if chosen and chosen != self.opencode_model: + _LOGGER.warning( + "[Agent] OpenCode's configured model is %s (the studio's " + "pick), but this backend is pinned to %s by " + "OPENCODE_MODEL. Unset that variable to follow the " + "picker.", chosen, self.opencode_model) + except Exception as exc: # noqa: BLE001 -- report the last known model + _LOGGER.debug("[Agent] current_model could not re-resolve: %s", exc) + return self.opencode_model or None def _opencode_base_url(self) -> str: """Same self-heal `OpenCodeChat._ensure_reachable` gives every chat diff --git a/weightslab/trainer/services/agent/opencode_chat.py b/weightslab/trainer/services/agent/opencode_chat.py index e67581a4..139f0abf 100644 --- a/weightslab/trainer/services/agent/opencode_chat.py +++ b/weightslab/trainer/services/agent/opencode_chat.py @@ -64,7 +64,7 @@ # unset -- which is exactly the "OpenCode picks WHATEVER model happens to be # configured, arbitrarily" failure this method exists to avoid in the first # place. -_DEFAULT_MODEL = "opencode/deepseek-v4-flash-free" +_DEFAULT_MODEL = "opencode/big-pickle" class OpenCodeError(RuntimeError): @@ -79,7 +79,7 @@ class OpenCodeChat: def __init__(self, base_url: str, model: Optional[str] = None, timeout: float = 60.0, workspace_dir: Optional[str] = None, url_is_explicit: bool = True, - model_is_explicit: bool = True): + model_is_explicit: bool = True, seed_model: Optional[str] = None): self.base_url = (base_url or "http://127.0.0.1:4096").rstrip("/") self.model = model self.timeout = timeout @@ -105,6 +105,13 @@ def __init__(self, base_url: str, model: Optional[str] = None, timeout: float = # (confirmed live: an image-generation preview model, useless for # this class's structured-JSON-reply use case). self.model_is_explicit = model_is_explicit + # A model from agent_config.yaml SEEDS the shared choice rather than + # pinning it: it is what to use when nobody has chosen anything yet + # (and it is then published, so the studio shows it), but a model + # picked in the UI afterwards wins. It used to pin, so a run started + # after picking a model in the studio quietly went back to the yaml + # value. OPENCODE_MODEL stays a hard pin -- automation needs one. + self.seed_model = (seed_model or "").strip() or None # -- wire helpers --------------------------------------------------- # @@ -317,7 +324,11 @@ def _ensure_model_resolved(self) -> None: this class's structured-JSON intent-parsing, since it isn't a text-reasoning model at all). + Returns a short source label ("pinned" | "opencode-config" | + "config-seed" | "default" | "kept") for the banner/logs. + Resolution order: + 0. `OPENCODE_MODEL` (model_is_explicit) -- a hard pin, left alone. 1. `GET /config`'s own `model` field -- the one the model picker writes back to opencode.json on every pick (opencodeClient.ts's setDefaultModel), so it's "whatever the user last actually @@ -337,22 +348,116 @@ def _ensure_model_resolved(self) -> None: explicitly chosen, land on the known-good free model" now means exactly that, with no provider-reported default able to override it. - Resolved once and cached on self.model. An explicit model - (model_is_explicit=True) is left alone -- deliberately chosen, not - a placeholder to override. + Re-checked before every turn, NOT cached for the life of the + process: the studio's model picker (.wl-ag-model) writes the pick + into OpenCode's own config via `PUT /config`, and a backend that + latched onto the model it saw at startup went on answering with the + old one for the rest of the run. `GET /config` is a local request on + the same machine, so following it per turn costs nothing next to the + completion it precedes. A failed read keeps the current model instead + of falling back. + + An explicit model (model_is_explicit=True -- OPENCODE_MODEL or + agent_config.yaml's `opencode_model`) is left alone: deliberately + chosen, not a placeholder to override, and NOT overridable from the + UI picker either. """ - if self.model_is_explicit or self.model: - return + if self.model_is_explicit: + return "pinned" + model_id = self._configured_model() + if model_id: + if model_id != self.model: + _LOGGER.info("[OpenCodeChat] following OpenCode's configured " + "model: %s (was %s)", model_id, self.model or "unset") + self.model = model_id + return "opencode-config" + # Nothing chosen anywhere yet -- fall back to the configured seed + # before the built-in default, so a project's agent_config.yaml still + # decides which model a fresh setup starts on. + if self.seed_model: + self.model = self.seed_model + return "config-seed" + # /config could not be read (or names no model): keep whatever was + # resolved on an earlier turn rather than dropping a working model for + # the fallback because one local request happened to fail. + if not self.model: + self.model = _DEFAULT_MODEL + return "default" + return "kept" + + def publish_model(self, model: Optional[str] = None) -> bool: + """Write `model` (default: the resolved one) into OpenCode's own config. + + Same call the studio's model picker makes, so whichever side starts + first leaves ONE answer behind for the other to read, and both ends of + a session agree on the model without talking to each other. + + GLOBAL scope, with the workspace route as fallback: verified live, + PATCH /config echoes the value back but does NOT change what GET + /config then reports, while PATCH /global/config + (`global.config.update`) does, immediately -- and GET /config is what + both sides read. + + Best-effort: a read-only config or an older server must never stop the + agent from working with the model it already resolved in-process. + """ + model = model or self.model + if not model: + return False + for path in ("/global/config", "/config"): + try: + with self._request(path, method="PATCH", body={"model": model}) as resp: + resp.read() + except Exception as exc: # noqa: BLE001 -- advisory write, never fatal + _LOGGER.debug("[OpenCodeChat] could not publish model %s via %s: %s", + model, path, exc) + continue + # Confirm against the effective config rather than trusting the + # echo: the workspace route answers 200 for a write it drops. + if self._configured_model() == model: + _LOGGER.info("[OpenCodeChat] published model %s to OpenCode (%s) at %s", + model, path, self.base_url) + return True + _LOGGER.debug("[OpenCodeChat] model %s could not be published to %s", + model, self.base_url) + return False + + def _configured_model(self) -> Optional[str]: + """The model OpenCode itself reports (GET /config) -- the shared choice + the studio picker writes and this backend follows.""" try: with self._request("/config") as resp: config = json.loads(resp.read().decode("utf-8")) - model_id = (config or {}).get("model") - if isinstance(model_id, str) and "/" in model_id: - self.model = model_id - return - except Exception: # noqa: BLE001 - fall through to the hardcoded default - pass - self.model = _DEFAULT_MODEL + except Exception: # noqa: BLE001 + return None + model_id = (config or {}).get("model") + return model_id if isinstance(model_id, str) and "/" in model_id else None + + def resolve_model(self, publish_default: bool = False): + """Resolve the model NOW instead of lazily on the first turn, and say + where it came from: ("pinned" | "pinned-published" | "opencode-config" + | "default" | "default-published" | "kept"). + + Called at agent start-up so the banner states the model actually in + use -- it used to print "(server default)" whenever nothing was pinned, + which read as "the studio's choice was ignored" even when the first + turn would have picked it up correctly. + + With publish_default=True, the model is also written back to + OpenCode's config whenever this side is the one deciding it -- a pinned + model (OPENCODE_MODEL), a seed from agent_config.yaml, or the built-in + fallback. A backend started BEFORE the studio then hands the UI the + model it is itself using, instead of the studio showing an unrelated + default while every backend query ran on the pinned one. + + A model that CAME from OpenCode's config is never re-published: there + is nothing to write, and doing so would fight the picker. + """ + self._ensure_reachable() + source = self._ensure_model_resolved() + if publish_default and source in ("default", "config-seed", "pinned") and self.publish_model(): + source = f"{source}-published" + return self.model, source def _call(self, prompt_value): from langchain_core.messages import AIMessage diff --git a/weightslab/trainer/services/data_service.py b/weightslab/trainer/services/data_service.py index dd065d33..a2b0b477 100755 --- a/weightslab/trainer/services/data_service.py +++ b/weightslab/trainer/services/data_service.py @@ -44,8 +44,8 @@ from weightslab.data import media_store from weightslab.trainer.trainer_tools import execute_df_operation, generate_overview, encode_image_to_raw_bytes from weightslab.data.data_utils import load_raw_image_array - # Image encoding / mask compression / proto helpers (extracted) + from weightslab.trainer.services.data_image_utils import ( rle_encode_mask, create_data_stat, @@ -146,6 +146,47 @@ def _media_chunk_bytes() -> int: _NON_MASK_TASKS = ("classification", "tabular") +def is_set_flag(value) -> bool: + """True only for a boolean column value that is really SET and true. + + Two traps, both hit in production, both of them "a missing value read as + set": + + * ``bool(float("nan")) is True`` in Python, and a nullable flag column is + full of NaN -- ``discarded`` for a sample nothing has written yet, or a + ``tag:*`` column for every sample that does not carry the tag. Reading it + with ``bool()`` / ``astype(bool)`` reported all of them as set: samples + greyed out as discarded while the dataframe said False, and every sample + wearing every tag. + * ``bool("False") is True`` as well, and a boolean column that has been + through the H5 store (where these columns become categorical) can come + back holding the STRINGS "True"/"False". + + So: missing is false, a string is read as a word, everything else falls + back to ``bool()``. + """ + try: + if value is None or pd.isna(value): + return False + except (TypeError, ValueError): + # pd.isna raises for some array-likes; those are not missing. + pass + if isinstance(value, str): + return value.strip().lower() in ("1", "true", "yes", "y", "t") + try: + return bool(value) + except Exception: # noqa: BLE001 -- an exotic value is not a set flag + return False + + +def set_flag_mask(series) -> "np.ndarray": + """``is_set_flag`` over a whole column, as a numpy bool array.""" + if series is None: + return np.zeros(0, dtype=bool) + return np.fromiter((is_set_flag(v) for v in series.tolist()), + dtype=bool, count=len(series)) + + def _is_non_mask_task(task_type) -> bool: """True when labels/predictions for this task must not be read as masks.""" return task_type in _NON_MASK_TASKS or is_generation_task(task_type) @@ -494,6 +535,23 @@ def rewrite_boolean_keywords_to_bitwise(code: str) -> str: return code +def _histogram_category_cap() -> int: + """Max distinct bars a categorical histogram returns (WL_HIST_CATEGORY_CAP). + + Beyond this the response is neither renderable nor informative -- the + remainder is folded into a single "(other)" bar. + """ + try: + return max(1, int(os.environ.get("WL_HIST_CATEGORY_CAP", "200"))) + except Exception: + return 200 + + +def _fast_view_enabled() -> bool: + """Differential view refresh. On by default; set WL_FAST_VIEW=0 to opt out.""" + return os.environ.get("WL_FAST_VIEW", "1") not in ("0", "false", "False") + + class DataService: """ @@ -1109,6 +1167,23 @@ def _pull_into_all_data_view_df(self): # merge + proxy conversion) a second time over the whole dataset every refresh. df = self._df_manager.get_collapse_annotations_to_samples_df(df) + # The collapse yields object dtype for signals//* columns. They are + # numeric by definition, and an object column sorts ~8x slower + # (6.17s vs 0.76s at 3.96M) while also slowing histogram binning, + # groupby and every differential write. Coerce them back here, at the + # single point the view is materialised. + for _c in df.columns: + if not str(_c).startswith("signals") or df[_c].dtype != object: + continue + try: + _num = pd.to_numeric(df[_c], errors="coerce") + # Only when nothing is lost: a genuinely non-numeric value + # means the column is not what we assume, so leave it be. + if _num.notna().sum() == df[_c].notna().sum(): + df[_c] = _num.astype("float32") + except Exception as _exc: + logger.debug("[DataService] dtype restore skipped for %r: %s", _c, _exc) + # Ensure sample_id is a column if it was the index df = safe_reset_index(df) @@ -1131,7 +1206,9 @@ def _pull_into_all_data_view_df(self): return df except Exception as e: - logger.debug(f"[DataService] Error pulling data view: {e}") + # Was debug: a swallowed failure here silently freezes the view at + # the previous snapshot, which looks exactly like "no new data". + logger.error("[DataService] Error pulling data view: %s", e, exc_info=True) # Use getattr to safely check for attribute during __init__ current_df = getattr(self, "_all_datasets_df", None) return current_df if current_df is not None else pd.DataFrame() @@ -1448,8 +1525,9 @@ def _compute_custom_signals(self): except Exception as e: logger.error(f"[DataService] Failed to compute signals for loader '{loader_name}': {e}") - # Force view update - self._slowUpdateInternals(force=True) + # Refresh signal values; differential unless the schema actually changed. + if not self._fastUpdateInternals(): + self._slowUpdateInternals(force=True) def _process_sample_row(self, args): """Process a single dataframe row to create a DataRecord.""" @@ -1477,6 +1555,8 @@ def _process_sample_row(self, args): skip_prediction_for_request = metadata_only_request # ====== Step 2: Load dataset lazily (avoid unnecessary IO for metadata-only) ====== + # Views the client explicitly asked for; empty => send everything. + _wanted_stats = set(getattr(request, "stats_to_retrieve", None) or []) needs_dataset = bool(request.include_raw_data) or (not skip_label_for_request) dataset = self._get_dataset(origin) if needs_dataset else None @@ -1569,10 +1649,10 @@ def _process_sample_row(self, args): # 'discarded' drives the grayed-out cell rendering, so it rides with the # image data as "1"/"0" (not treated as analytical metadata). This keeps # the gray-out reliable on every grid (re)fetch / scroll. - try: - _discarded_str = "1" if bool(row.get(SampleStatsEx.DISCARDED.value)) else "0" - except Exception: - _discarded_str = "0" + # is_set_flag, not bool(): a nullable flag full of NaN read as + # set is what greyed out every sample the model had seen. + _discarded_str = "1" if is_set_flag( + row.get(SampleStatsEx.DISCARDED.value)) else "0" data_stats.append( create_data_stat( SampleStatsEx.DISCARDED.value, 'string', shape=[1], value_string=_discarded_str, thumbnail=b"" @@ -2150,14 +2230,91 @@ def _json_default(o): target_height=target_height, ) - data_stats.append( - create_data_stat( - name='raw_data', - stat_type='bytes', - thumbnail=raw_data_bytes, - shape=raw_shape, + # An explicit stats_to_retrieve means the client knows which + # views it will draw; raw_data duplicates view rank 0, so + # only send it when actually asked for. + if not _wanted_stats or 'raw_data' in _wanted_stats: + data_stats.append( + create_data_stat( + name='raw_data', + stat_type='bytes', + thumbnail=raw_data_bytes, + shape=raw_shape, + ) ) - ) + + # Paired/multi-view datasets (e.g. a source+edited image + # pair) can optionally expose additional named views via + # extra_images() -- send each as its own 'image_' + # stat so the frontend renders it as its own grid column + # (isImageFieldName() already recognizes 'image_*'). This + # duck-typed hook is a no-op for datasets that don't + # define it. raw_data above already covers view rank 0 + # (e.g. 'source'); extra_images() may repeat that view + # under its own name too -- one small duplicated + # thumbnail, traded for not having to assume which named + # view is redundant across arbitrary datasets. + # Probe the UNWRAPPED dataset: `dataset` is WL's tracking + # wrapper and does not forward extra_images, so testing it + # silently disables every named view. + if hasattr(ds, "extra_images"): + try: + extra_views = ds.extra_images(ds_idx) or {} + except Exception as e: + extra_views = {} + logger.debug(f"extra_images failed for sample_id={sample_id}: {e}") + for view_name, view_pil in extra_views.items(): + if view_pil is None: + continue + # Filter BEFORE resize/encode -- that is the cost. + # Still ADVERTISE the view with an empty thumbnail: + # the panel builds its modality list from the stats + # present, so omitting it entirely would delete the + # toggle and make the view unrecoverable. + if (_wanted_stats + and ("image_%s" % view_name) not in _wanted_stats): + data_stats.append( + create_data_stat( + name="image_%s" % view_name, + stat_type='bytes', + thumbnail=b"", + shape=[], + ) + ) + continue + try: + resized_view = view_pil + if resized_view.size != (target_width, target_height): + _view_resample = ( + Image.Resampling.LANCZOS if is_full_resolution + else Image.Resampling.BILINEAR + ) + resized_view = resized_view.resize( + (target_width, target_height), _view_resample + ) + view_bytes, view_shape, _ = encode_image_to_raw_bytes( + np_img=None, + middle_pil=resized_view, + original_shape=[], + is_volumetric=False, + is_full_resolution=is_full_resolution, + target_width=target_width, + target_height=target_height, + ) + data_stats.append( + create_data_stat( + name=f"image_{view_name}", + stat_type='bytes', + thumbnail=view_bytes, + shape=view_shape, + ) + ) + del view_bytes, resized_view + except Exception as e: + logger.debug( + f"extra_images encode failed for sample_id={sample_id} " + f"view={view_name}: {e}" + ) # For video samples the bytes above are only the poster # frame, so advertise the clip's shape here. This lets the @@ -2470,16 +2627,49 @@ def _sample_id_sortable_series(self, values): return numeric return values.astype(str) + def _numeric_like_sort_cols(self, df: pd.DataFrame, by) -> set: + """Sort columns whose values are strings but mean numbers. + + '1916469' < '191647' lexicographically but not numerically, so sorting a + numeric-valued string column by raw string order is simply wrong. Only + object/string columns are candidates; genuinely numeric dtypes already + sort correctly and must not be touched. + """ + out = set() + by_list = [by] if isinstance(by, str) else list(by or []) + for col in by_list: + if col == SampleStatsEx.SAMPLE_ID.value: + out.add(col) + continue + try: + s = df[col] if col in df.columns else None + if s is None or pd.api.types.is_numeric_dtype(s) or hasattr(s, "cat"): + continue + probe = s.dropna() + if probe.empty: + continue + # Cheap decision on a sample: a full 4M-row coercion here would + # cost more than the sort it is meant to correct. + head = probe.head(2048) + coerced = pd.to_numeric(head, errors="coerce") + if coerced.notna().all(): + out.add(col) + except Exception: + continue + return out + def _sort_values_numeric_aware(self, df: pd.DataFrame, sort_params: dict) -> None: - """Sort dataframe while treating sample_id as numeric when possible.""" + """Sort dataframe, ordering numeric-valued string columns numerically.""" params = dict(sort_params) - if params.get("key") is None and self._sort_includes_sample_id(params.get("by")): - def _key(series: pd.Series): - if str(getattr(series, "name", "")) == SampleStatsEx.SAMPLE_ID.value: - return self._sample_id_sortable_series(series) - return series + if params.get("key") is None: + numeric_like = self._numeric_like_sort_cols(df, params.get("by")) + if numeric_like: + def _key(series: pd.Series): + if str(getattr(series, "name", "")) in numeric_like: + return self._sample_id_sortable_series(series) + return series - params["key"] = _key + params["key"] = _key df.sort_values(inplace=True, **params) @@ -3527,6 +3717,38 @@ def _apply_agent_operation(self, df, func: str, params: dict) -> str: # silently stops refreshing). orig_index_names = [n for n in df.index.names if n is not None] + # Fast path: nothing in `by` is an index level, so the frame can be + # sorted where it stands. Avoids reset_index + astype(int)/astype(str) + # over every sample_id + a set_index that re-factorizes 3.96M string + # keys -- measured as the bulk of a ~20s sort. + _fp_by = params.get("by") + _fp_list = [_fp_by] if isinstance(_fp_by, str) else list(_fp_by or []) + _fp_res = [ + (c if c in df.columns + else ("signals//" + c if ("signals//" + c) in df.columns else c)) + for c in _fp_list + ] + if (_fp_res + and all(c in df.columns for c in _fp_res) + and not any(c in orig_index_names for c in _fp_list) + and not any(c in orig_index_names for c in _fp_res)): + _fp_params = dict(params) + _fp_params["by"] = (_fp_res if isinstance(_fp_by, (list, tuple)) + else _fp_res[0]) + try: + # Through the helper, NOT df.sort_values directly: a + # numeric-valued string column (group_id, target, ...) + # otherwise sorts lexicographically -- '1916469' before + # '191647'. + self._sort_values_numeric_aware(df, _fp_params) + return "Applied operation: sort_values" + except (TypeError, ValueError, KeyError) as _fp_exc: + # Mixed dtypes or an unexpected key: fall through to the + # original reset/restore path rather than failing the query. + logger.debug( + "[sort] fast path declined (%s); using index round-trip", + type(_fp_exc).__name__) + def _restore_index(): cols = [n for n in orig_index_names if n in df.columns] if cols and not isinstance(df.index, pd.MultiIndex): @@ -3789,12 +4011,137 @@ def _bg_view_refresh(self) -> None: real rebuild+swap via force=True OFF the request path, then releases the guard so a later stale read can trigger another. Never raises into a request.""" try: - self._slowUpdateInternals(force=True) + if not self._fastUpdateInternals(): + self._slowUpdateInternals(force=True) except Exception: logger.exception("[ViewRefresh] background view refresh failed") finally: self._refresh_in_flight.release() + # Columns the trainer mutates. Structural columns (origin, edit_prompt, + # task_type, ...) never change after registration, so a differential sync + # only has to carry these. + _FAST_SYNC_PREFIXES = ("signals", "last_seen", "discarded", "prediction", "target") + + def _fast_sync_columns(self, view): + return [c for c in view.columns + if str(c).startswith(self._FAST_SYNC_PREFIXES)] + + + + def _fastUpdateInternals(self, max_dirty: int = 250_000) -> bool: + """O(change) view refresh. True if applied, False -> caller must rebuild. + + Falls back when there is no view yet, when a dirty row is absent from + the view (new rows => structural change), or when the backlog is large + enough that a rebuild is cheaper. + """ + if not _fast_view_enabled(): + return False # opt-out: behave exactly as before + view = self._all_datasets_df + dfm = self._df_manager + if view is None or getattr(view, "empty", True) or dfm is None: + return False + # A manager without dirty tracking cannot serve a delta -- rebuild. + if not (hasattr(dfm, "take_view_dirty") and hasattr(dfm, "get_source_rows")): + return False + # A column the ledger has but the view lacks can only arrive via a full + # rebuild -- the differential write below addresses existing columns + # only. Per-sample signal columns are created on their first write, so + # on a fresh ledger the view predates them and would never gain them. + try: + _src = getattr(dfm, "_df", None) + if _src is not None: + _have = set(view.columns) + _missing = [c for c in _src.columns + if str(c).startswith(self._FAST_SYNC_PREFIXES) + and c not in _have] + if _missing: + return False + except Exception: + pass + + dirty = dfm.take_view_dirty(limit=max_dirty) + if dirty is None: + return False # backlog too large; rebuild is cheaper + if not dirty: + return True # nothing changed since last sync + + sids = [str(s) for s in dirty] + # Address rows by LABEL. pandas keeps a hash engine on the index, built + # in C and cached, so this needs no precomputed position map -- and a + # label that is absent surfaces below as a no-match rather than as a + # silently wrong row. + keep = sids + + cols = self._fast_sync_columns(view) + if not cols: + return True + sub = dfm.get_source_rows(keep, columns=[c for c in cols if c in view.columns]) + if sub is None or sub.empty: + return True + # Collapse the source's per-annotation rows to ONE row per sample the + # same way the view itself was built (see + # get_collapse_annotations_to_samples_df): the canonical row is + # annotation_id == 0, and sample-level columns live only there. + # + # This used to keep the LAST annotation row, which for a multi-instance + # sample carries NaN in every sample-level column -- so each sample the + # trainer touched had its view row's `discarded` / `prediction` / + # `target` overwritten with NaN. And `bool(float("nan"))` is True in + # Python, so GetDataSamples then reported discarded="1" and the studio + # greyed the sample out, progressively, exactly as the model worked + # through the dataset -- while the dataframe itself still said False. + if isinstance(sub.index, pd.MultiIndex): + ANNOT = SampleStatsEx.INSTANCE_ID.value + names = list(getattr(sub.index, "names", []) or []) + if ANNOT in names: + annot = sub.index.get_level_values(ANNOT) + try: + canonical = np.asarray(annot).astype(int) == 0 + except (TypeError, ValueError): + canonical = np.array([str(a) in ("0", "0.0") for a in annot]) + if canonical.any(): + sub = sub[canonical] + sub = sub.droplevel(-1) + # Whatever is left, one row per sample: prefer the FIRST (the canonical + # row when the level was present, the first occurrence otherwise) -- + # never the last, for the reason above. + sub = sub[~sub.index.duplicated(keep="first")] + + # Only rows the view actually holds; a structural change (new sample) + # must still fall back to the full rebuild rather than be invented here. + SID = SampleStatsEx.SAMPLE_ID.value + _names = list(getattr(view.index, "names", []) or []) + view_keys = (view.index.get_level_values(SID) + if isinstance(view.index, pd.MultiIndex) and SID in _names + else view.index) + # Looked up the other way round -- `sub` (deduplicated just above, so + # unique) is the index being searched, and the VIEW's keys are the + # target. Searching the view's keys instead raised + # InvalidIndexError("Reindexing only valid with uniquely valued Index + # objects") whenever one sample_id appeared under two origins, which + # the view's own (origin, sample_id) index exists precisely to allow -- + # and the differential refresh then failed every time, silently falling + # back to the full rebuild. This direction also updates BOTH rows of + # such a sample, which is the only thing the source (indexed by + # sample_id alone) can mean. + _view_keys = pd.Index(view_keys.astype(str)) + _sub_keys = pd.Index(sub.index.astype(str)) + _src = _sub_keys.get_indexer(_view_keys) # sub row per view row, -1 if none + _rows = np.flatnonzero(_src >= 0) + if _rows.size == 0: + return True + # A dirty sample the view does not hold is a structural change (a new + # sample): only the full rebuild can add it. + if len(set(_sub_keys)) != len(set(_view_keys[_rows])): + return False + _take = _src[_rows] + for c in sub.columns: + _ci = view.columns.get_loc(c) + view.iloc[_rows, _ci] = sub[c].to_numpy()[_take] + return True + def _slowUpdateInternals(self, force: bool = False, reset_view: bool = False) -> None: """Update the internal dataframe view with the latest data from the manager. @@ -3965,6 +4312,13 @@ def _slowUpdateInternals(self, force: bool = False, reset_view: bool = False) -> # Atomic swap to make the new view available to readers self._all_datasets_df = updated_df self._last_internals_update_time = current_time + # The rebuilt view reflects every row, so the differential backlog is + # satisfied. This is the only point at which that is true. + try: + if self._df_manager is not None: + self._df_manager.clear_view_dirty() + except Exception: + pass finally: held_ms = (time.time() - t_held_start) * 1000 @@ -4180,8 +4534,11 @@ def _build_metadata_only_response(self, df_slice: pd.DataFrame, requested_cols=N _DataStat(name=col, type="string", shape=[1], value_string=v) ) else: - # Boolean tag: presence indicator "1" when True. - bools = series.astype(bool).tolist() + # Boolean tag: presence indicator "1" when True. is_set_flag, + # not astype(bool): a tag column is NaN for every sample that + # does not carry the tag, and astype(bool) turns NaN into True + # -- which showed every tag on every sample. + bools = set_flag_mask(series) for i, b in enumerate(bools): if b: row_stats[i].append( @@ -4482,7 +4839,8 @@ def _process_get_data_samples(self, request, context): ) # Trigger update if needed (it has its own internal locking) - self._slowUpdateInternals() + if not self._fastUpdateInternals(): + self._slowUpdateInternals() # Atomic snapshot of the current authoritative dataframe current_df = self._all_datasets_df @@ -4793,37 +5151,39 @@ def ApplyDataQuery(self, request, context): operations = self._parse_direct_query(request.query) # Apply operations with lock - with self._watched_lock("_lock[ApplyDataQuery/ops]"): - # Skip the forced full-view rebuild for SORT-ONLY operations. Sorting just - # re-orders the existing snapshot, so a fresh collapse+combine (hundreds of - # ms on large views, and — being lock-held — contends with the training - # thread for multi-second stalls) is unnecessary. Filters/edits still refresh - # so they operate on the latest data. The view is frozen on direct queries - # anyway (_is_filtered=True), so it wasn't auto-refreshing mid-sort regardless. - _SORT_FUNCS = {"df.sort_values", "df.sort_index", "df.sort_view_slice"} - is_sort_only = bool(operations) and all( - op.get("function") in _SORT_FUNCS for op in operations) - if not is_sort_only: - self._slowUpdateInternals(force=True) # Refresh internals before applying non-sort operations - - # Work on a copy to allow concurrent readers to see a consistent state - df = self._all_datasets_df # Remove copy because memory waste and slowdown - messages = [] + _SORT_FUNCS = {"df.sort_values", "df.sort_index", "df.sort_view_slice"} + is_sort_only = bool(operations) and all( + op.get("function") in _SORT_FUNCS for op in operations) + def _run_ops(target): + out = [] for op in operations: - func = op.get("function") - params = op.get("params", {}) or {} - msg = self._apply_agent_operation(df, func, params) - messages.append(msg) - - final_message = " | ".join(messages) if messages else "No operation performed" - - # Atomic swap - self._all_datasets_df = df - - # Direct queries are manipulations -> Freeze the view - if operations: - self._is_filtered = True + out.append(self._apply_agent_operation( + target, op.get("function"), op.get("params", {}) or {})) + return " | ".join(out) if out else "No operation performed" + + if is_sort_only: + # Sorting 4M rows costs ~10s; under _lock that stalls the trainer + # for the whole duration (measured 11716ms). Sort a shallow copy + # off-lock, then hold the lock only for the reference swap. + base = self._all_datasets_df + df = base.copy(deep=False) if base is not None else base + final_message = _run_ops(df) + # No position map to rebuild: the differential refresh addresses + # rows by label through pandas' own index engine, so reordering + # the view invalidates nothing. + with self._watched_lock("_lock[ApplyDataQuery/swap]"): + self._all_datasets_df = df + if operations: + self._is_filtered = True + else: + with self._watched_lock("_lock[ApplyDataQuery/ops]"): + self._slowUpdateInternals(force=True) + df = self._all_datasets_df + final_message = _run_ops(df) + self._all_datasets_df = df + if operations: + self._is_filtered = True return self._build_success_response( df=df, @@ -5047,17 +5407,38 @@ def GetHistogram(self, request, context): if df is None or df.empty: return pb2.HistogramResponse( success=False, message="empty dataframe view", total_rows=0, bins=[]) - df = safe_reset_index(df) + # safe_reset_index copies AND consolidates the entire frame (~70% of + # this RPC at 3.96M x 19 by py-spy). It is only needed to reach fields + # that live in the index; when they are already columns, use the frame + # as it stands. + def _field(frame, name): + """Series for *name* whether it is a column or an index level.""" + if name in frame.columns: + return frame[name] + names = list(getattr(frame.index, "names", []) or []) + if name in names: + return pd.Series(frame.index.get_level_values(name), + index=frame.index) + if getattr(frame.index, "name", None) == name: + return pd.Series(frame.index, index=frame.index) + return None + + # Only reset when the histogrammed column itself cannot be reached. + # safe_reset_index copies AND block-consolidates the whole frame + # (~82% of this RPC by py-spy); get_level_values is ~0.02s. + if _field(df, column) is None: + df = safe_reset_index(df) n = len(df) if column not in df.columns: return pb2.HistogramResponse( success=False, message=f"column '{column}' not in view", total_rows=n, bins=[]) - origin = (df["origin"].astype(str).to_numpy() if "origin" in df.columns - else np.full(n, "")) - disc = (df["discarded"].astype(bool).to_numpy() if "discarded" in df.columns - else np.zeros(n, bool)) + _o = _field(df, "origin") + origin = _o.astype(str).to_numpy() if _o is not None else np.full(n, "") + _d = _field(df, "discarded") + # Same trap as above: NaN in a nullable flag is NOT "discarded". + disc = set_flag_mask(_d) if _d is not None else np.zeros(n, bool) # Detect whether column is categorical (string/object) or numeric. # A column is numeric if ANY value coerces to a finite number — even @@ -5068,7 +5449,9 @@ def GetHistogram(self, request, context): # as a spurious "unset" bar). We therefore treat as categorical only # a genuine pandas ``category`` dtype, or a column whose values do # not coerce to any numeric value at all (pure strings). - col_series = df[column] + col_series = _field(df, column) + if col_series is None: + col_series = df[column] numeric_vals = pd.to_numeric(col_series, errors="coerce") is_category_dtype = ( str(col_series.dtype) == "category" or hasattr(col_series, "cat") @@ -5087,20 +5470,44 @@ def GetHistogram(self, request, context): if is_categorical: # --- Categorical path --- labels = col_series.astype(str).where(col_series.notna(), "") - gf = pd.DataFrame({"l": labels, "o": origin, "d": disc}) - total_count = gf.groupby("l")["l"].count().rename("count") + # Count first with a single-key value_counts, then restrict the + # three-key breakdown to the rows that survive the cap. Grouping + # all 3.96M rows by (label, discarded, origin) when the column is + # free text builds 722,870 groups to then discard all but 200. + total_count = labels.value_counts().rename("count") + _cap_pre = _histogram_category_cap() + _keep = set(total_count.iloc[:_cap_pre].index) + _m = labels.isin(_keep).to_numpy() sub_map: dict = {} - for (lbl, d, o), c in gf.groupby(["l", "d", "o"]).size().items(): - sub_map.setdefault(str(lbl), []).append( - pb2.HistogramSubBar(origin=str(o), discarded=bool(d), count=int(c))) + if _m.any(): + gf = pd.DataFrame({"l": labels.to_numpy()[_m], + "o": origin[_m], "d": disc[_m]}) + for (lbl, d, o), c in gf.groupby(["l", "d", "o"]).size().items(): + sub_map.setdefault(str(lbl), []).append( + pb2.HistogramSubBar(origin=str(o), discarded=bool(d), count=int(c))) + # Cap the output: a free-text column (e.g. edit_prompt) has one + # category per sample -- 722,870 bars / 54 MB / 30.5s measured, + # which no viewer can draw. Keep the top-N by count and fold the + # remainder into one "(other)" bar so the response stays bounded. + _ordered = total_count.sort_values(ascending=False) + _cap = _histogram_category_cap() + _head, _tail = _ordered.iloc[:_cap], _ordered.iloc[_cap:] cat_bars = [ pb2.CategoricalHistogramBar( label=str(lbl), count=int(cnt), sub_bars=sub_map.get(str(lbl), []), ) - for lbl, cnt in total_count.sort_values(ascending=False).items() + for lbl, cnt in _head.items() ] + if len(_tail): + cat_bars.append(pb2.CategoricalHistogramBar( + label="(other: %d categories)" % len(_tail), + count=int(_tail.sum()), + sub_bars=[], + )) + logger.info("[HistCat] column=%s capped %d categories -> %d bars", + column, len(_ordered), len(cat_bars)) logger.info("[HistCat] column=%s rows=%d categories=%d", column, n, len(cat_bars)) return pb2.HistogramResponse( @@ -5113,6 +5520,11 @@ def GetHistogram(self, request, context): ) # --- Numeric path (unchanged) --- + # Each bar covers a fixed slice of the VIEW: total_rows / max_bins + # samples. That is what makes the chart show density -- a column only + # 0.2% populated shows a few filled bars and the rest empty. Binning + # over just the rows that carry a value makes the chart look equally + # full at any coverage, which reads as "everything has a value". bars = max(1, min(n, max_bins)) vals = numeric_vals.to_numpy() edges = (np.arange(bars + 1) * n) // bars @@ -5627,7 +6039,8 @@ def EditDataSample(self, request, context): with self._watched_lock("_lock[EditDataSample/__copy_metadata__]"): try: - self._slowUpdateInternals() + if not self._fastUpdateInternals(): + self._slowUpdateInternals() if self._all_datasets_df is None or self._all_datasets_df.empty: return pb2.DataEditsResponse( success=False, @@ -5795,7 +6208,8 @@ def EditDataSample(self, request, context): with self._watched_lock("_lock[EditDataSample/delete-col]"): try: - self._slowUpdateInternals() + if not self._fastUpdateInternals(): + self._slowUpdateInternals() if self._all_datasets_df is None or self._all_datasets_df.empty: return pb2.DataEditsResponse( success=False, @@ -5833,7 +6247,8 @@ def EditDataSample(self, request, context): # Kick a background view-refresh (non-blocking) — the in-memory view # is already consistent after the drop above, so blocking inline rebuild # is unnecessary and causes the gRPC response to stall for 5-10 s. - self._slowUpdateInternals() + if not self._fastUpdateInternals(): + self._slowUpdateInternals() return pb2.DataEditsResponse( success=True, @@ -5856,7 +6271,8 @@ def EditDataSample(self, request, context): with self._watched_lock("_lock[EditDataSample/__discard_by_tag__]"): try: - self._slowUpdateInternals() + if not self._fastUpdateInternals(): + self._slowUpdateInternals() if self._all_datasets_df is None or self._all_datasets_df.empty: return pb2.DataEditsResponse( success=False, @@ -5865,7 +6281,8 @@ def EditDataSample(self, request, context): tag_col = f"{SampleStatsEx.TAG.value}:{tag_name}" if tag_col not in self._all_datasets_df.columns: - self._slowUpdateInternals() + if not self._fastUpdateInternals(): + self._slowUpdateInternals() df = safe_reset_index(self._all_datasets_df) if tag_col not in df.columns: return pb2.DataEditsResponse( @@ -6079,7 +6496,8 @@ def GetDataSplits(self, request, context): # IMPORTANT: keep lock ordering consistent (_update_lock -> _lock). # Calling _slowUpdateInternals() while holding _lock can deadlock # with concurrent readers/writers under high UI refresh pressure. - self._slowUpdateInternals() + if not self._fastUpdateInternals(): + self._slowUpdateInternals() if context is not None and not context.is_active(): return pb2.DataSplitsResponse(success=False, split_names=[]) diff --git a/weightslab/trainer/trainer_tools.py b/weightslab/trainer/trainer_tools.py index 136dd34c..1b2d467c 100644 --- a/weightslab/trainer/trainer_tools.py +++ b/weightslab/trainer/trainer_tools.py @@ -434,10 +434,15 @@ def _get_input_tensor_for_sample(dataset, sample_id, device): def process_sample(sid, dataset, do_resize, resize_dims, experiment): try: - if hasattr(dataset, "_getitem_raw"): - tensor, idx, label = dataset._getitem_raw(id=sid) - else: - tensor, idx, label = dataset[sid] + # _getitem_raw returns (data, id, target, *metadata) by contract -- datasets + # implementing get_items() with metadata yield 4+ elements, so a fixed + # 3-way unpack raises "too many values to unpack" and kills every + # thumbnail. Unpack positionally instead. + _res = dataset._getitem_raw(id=sid) if hasattr(dataset, "_getitem_raw") else dataset[sid] + if not isinstance(_res, (tuple, list)): + _res = (_res, sid, None) + tensor = _res[0] + label = _res[2] if len(_res) > 2 else None if isinstance(tensor, torch.Tensor): img = tensor.detach().cpu() diff --git a/weightslab/ui/server.py b/weightslab/ui/server.py index 53c7eec7..0d7fed6b 100644 --- a/weightslab/ui/server.py +++ b/weightslab/ui/server.py @@ -1577,6 +1577,9 @@ class _UIRequestHandler(BaseHTTPRequestHandler): grpc_auth_token: Optional[str] = None rpc_timeout: float = 300.0 experiment_dir: Optional[str] = None + # The gRPC port this server proxies to; used to pick the right backend's + # experiment directory out of the active-experiment marker. + backend_port: Optional[int] = None # -- logging: quiet by default, honour WEIGHTSLAB_UI_VERBOSE ------------- # def log_message(self, fmt, *args): # noqa: D401 @@ -1612,6 +1615,9 @@ def do_GET(self): # noqa: N802 if path == "/agent-server/status": self._send_json(HTTPStatus.OK, _opencode_session.status()) return + if path == "/agent-server/model": + self._get_shared_model() + return if path == "/agent-server/loop/list": self._send_json(HTTPStatus.OK, {"loops": _loop_registry.list()}) return @@ -1645,6 +1651,8 @@ def do_POST(self): # noqa: N802 self._start_local_notebook() elif path == "/agent-server/start": self._start_agent_server() + elif path == "/agent-server/model": + self._set_shared_model() elif path == "/agent-server/loop/start": self._start_loop() elif path.startswith("/agent-server/loop/") and path.endswith("/stop"): @@ -1745,9 +1753,120 @@ def _collect_metadata(self): # ------------------------------------------------------------------ # # Local Jupyter Notebook launcher (landing-page button) # ------------------------------------------------------------------ # + # ------------------------------------------------------------------ # + # Shared agent model, proxied SAME-ORIGIN + # ------------------------------------------------------------------ # + # OpenCode's own config holds the one model every client of that server + # agrees on: the studio's picker, the OpenCode CLI, `weightslab agent + # model`, and the backend SDK agent (which reads it to choose the model for + # its own queries). The browser could call OpenCode directly -- but only + # while OpenCode's --cors allowlist happens to contain the exact origin the + # page is served from, which quietly fails for a LAN address, a tunnel + # hostname, or an OpenCode somebody started by hand with no --cors at all. + # The page then showed a model nobody else was using, and picking one + # changed nothing outside the tab. + # + # Same-origin here means no preflight and no allowlist: this server talks + # to OpenCode over plain HTTP on the machine they share. + def _opencode_base_url(self) -> Optional[str]: + status = _opencode_session.status() + url = status.get("url") if isinstance(status, dict) else None + if isinstance(url, str) and url.strip(): + return url.rstrip("/") + env_url = (os.environ.get("OPENCODE_URL") or "").strip() + return env_url.rstrip("/") if env_url else None + + def _opencode_json(self, path: str, method: str = "GET", body: Optional[dict] = None, + timeout: float = 10.0): + """One request to the local OpenCode server; (status, parsed-json-or-None).""" + base = self._opencode_base_url() + if not base: + return None, None + import urllib.error + import urllib.request + data = json.dumps(body).encode("utf-8") if body is not None else None + headers = {"Content-Type": "application/json"} if data is not None else {} + req = urllib.request.Request(base + path, data=data, headers=headers, method=method) + try: + with urllib.request.urlopen(req, timeout=timeout) as resp: + raw = resp.read().decode("utf-8", "replace") + try: + return resp.status, json.loads(raw or "null") + except ValueError: + return resp.status, None + except urllib.error.HTTPError as exc: + return exc.code, None + except Exception as exc: # noqa: BLE001 + logger.debug("[ui] OpenCode %s %s failed: %s", method, path, exc) + return None, None + + def _get_shared_model(self): + status, payload = self._opencode_json("/config") + # A transport failure AND an error status both mean "we do not know + # which model is configured" -- reporting ok:True with model=None there + # would tell the page "nothing is configured", and it would then show + # its own default over whatever the backend is really using. + if status is None or not (200 <= int(status) < 300): + self._send_json(HTTPStatus.OK, {"ok": False, "model": None, + "error": "OpenCode server is not reachable"}) + return + model = (payload or {}).get("model") if isinstance(payload, dict) else None + if not (isinstance(model, str) and "/" in model): + model = None + self._send_json(HTTPStatus.OK, {"ok": True, "model": model}) + + def _set_shared_model(self): + body = self._read_json_body() + model = str((body or {}).get("model") or "").strip() + if "/" not in model: + self._send_json(HTTPStatus.BAD_REQUEST, + {"ok": False, "error": "model must be \"providerID/modelID\""}) + return + # Global scope first: PATCH /config (workspace scope) answers 200 and + # echoes the value back without changing what GET /config reports. + for target in ("/global/config", "/config"): + status, _ = self._opencode_json(target, method="PATCH", body={"model": model}) + if status is None or not (200 <= int(status) < 300): + continue + _, payload = self._opencode_json("/config") + if isinstance(payload, dict) and payload.get("model") == model: + self._send_json(HTTPStatus.OK, {"ok": True, "model": model, "via": target}) + return + self._send_json(HTTPStatus.OK, { + "ok": False, "model": None, + "error": "OpenCode did not accept the model (config may be read-only)"}) + + def _experiment_dir_path(self) -> str: + """The experiment directory to browse for this run's own files. + + A RUNNING training backend's own resolved root_log_dir comes first: it + is where reports and notebooks are actually written, and it is not + always the directory this UI established -- training may have been + pointed elsewhere by a config file's `root_log_dir:`, or (before the + marker existed) have fallen through to a temp directory. Listing this + server's own directory then showed an empty reports/ right after a + report had been generated. + + Only a LIVE backend counts: the marker outlives the process that wrote + it, and a finished run's directory must never hijack the listing of a + UI that was given an experiment directory of its own. Falls back to the + UI's own directory, the environment, then the working directory. + """ + try: + from weightslab.utils.active_experiment import live_backend_experiment_dir + # By PORT: with two experiments up, "the live backend" is ambiguous + # and picking the wrong one shows the other experiment's reports. + backend_dir = live_backend_experiment_dir( + getattr(self, "backend_port", None)) + except Exception: # noqa: BLE001 -- never break a listing on the marker + backend_dir = None + return (backend_dir + or self.experiment_dir + or os.environ.get("WEIGHTSLAB_ROOT_LOG_DIR") + or os.getcwd()) + def _notebooks_dir_path(self) -> str: - experiment_dir = self.experiment_dir or os.environ.get("WEIGHTSLAB_ROOT_LOG_DIR") or os.getcwd() - return os.path.join(experiment_dir, "notebooks") + return os.path.join(self._experiment_dir_path(), "notebooks") def _read_json_body(self) -> dict: try: @@ -2088,13 +2207,12 @@ def _start_local_notebook(self): # (ApplyDataQuery -> the "generate_experiment_report" action, see # data_service.py) -- these two endpoints only browse what's already on # disk under /reports/, exactly like the local-notebook - # endpoints above browse /notebooks/. Same assumption: - # this UI server and the connected training backend share a filesystem - # (the documented `weightslab start` usage), so root_log_dir resolved - # here is the same directory the backend wrote reports into. + # endpoints above browse /notebooks/. The one assumption + # left is a shared filesystem (the documented `weightslab start` usage): + # WHICH directory is resolved by _experiment_dir_path(), which prefers the + # backend's own recorded root_log_dir over this server's. def _agent_history_dir_path(self) -> str: - experiment_dir = self.experiment_dir or os.environ.get("WEIGHTSLAB_ROOT_LOG_DIR") or os.getcwd() - return os.path.join(experiment_dir, "agent") + return os.path.join(self._experiment_dir_path(), "agent") def _dump_agent_history(self): """Write the agent conversation to the experiment directory. @@ -2138,8 +2256,7 @@ def _dump_agent_history(self): self._send_json(HTTPStatus.OK, {"ok": True, "path": path}) def _reports_dir_path(self) -> str: - experiment_dir = self.experiment_dir or os.environ.get("WEIGHTSLAB_ROOT_LOG_DIR") or os.getcwd() - return os.path.join(experiment_dir, "reports") + return os.path.join(self._experiment_dir_path(), "reports") def _list_experiment_reports(self): reports_dir = self._reports_dir_path() @@ -2456,6 +2573,7 @@ def serve_ui( { "static_root": root, "channel": channel, + "backend_port": backend_port, "api_prefix": "/api", "grpc_auth_token": grpc_auth_token, "experiment_dir": experiment_dir, diff --git a/weightslab/utils/active_experiment.py b/weightslab/utils/active_experiment.py new file mode 100644 index 00000000..fb53f9b2 --- /dev/null +++ b/weightslab/utils/active_experiment.py @@ -0,0 +1,340 @@ +"""Cross-process handoff of the active experiment directory. + +``weightslab start`` establishes the experiment directory (checkpoints, logs, +``notebooks/``, ``reports/``) and exports ``WEIGHTSLAB_ROOT_LOG_DIR`` -- but it +can only export it into ITS OWN process. A training run launched from a second +terminal, or by ``weightslab start example``, is a different process tree and +never saw that variable: it fell through to a throwaway ``tempfile.mkdtemp()``, +so the run wrote into ``%TEMP%\\tmpXXXXXXXX`` while the UI listed an empty +``reports/`` (the "right-click Generate Report shows nothing, yet I generated +reports" symptom) from the directory it had established itself. + +Hence this marker file: one small JSON document, per user, that both sides +write to and read from. Each side is a LIST, because running two experiments +side by side -- a classification UI on one port, a segmentation UI on another -- +is a supported thing to do. With a single slot per side, the second +``weightslab start`` erased the first, and both UIs would then have listed the +reports of whichever backend started last. + + { + "ui": [{"root_log_dir": "...", "pid": 123, "ui_port": 8080, + "backend_port": 50051, "updated_at": "..."}, ...], + "backend": [{"root_log_dir": "...", "pid": 456, "grpc_port": 50051, + "updated_at": "..."}, ...] + } + +* ``ui`` is written by ``weightslab start`` -- the directory it established. + A later training process with nothing configured adopts it, which is what + makes the two halves land in the same experiment. +* ``backend`` is written by ``wl.serve()`` -- the directory training ACTUALLY + resolved, whatever the route (an explicit ``root_log_dir:`` in a config file, + the environment, or the ``ui`` value above). The UI prefers it when listing + reports and notebooks, so those lists stay right even when training was + pointed somewhere the UI never chose. + +Entries are keyed by pid: a process replaces its own, dead ones are pruned on +every write, and the list is bounded. Readers never guess -- a port pins the +entry when the caller knows one (a UI asking for ITS backend), a single live +entry is unambiguous, and two or more return nothing with a line in the log, +leaving the caller its own directory. Showing one experiment's reports inside +another experiment's UI, or writing a run into the wrong experiment, is worse +than declining to answer. + +Neither side is required: every reader validates that the recorded directory +still exists and falls back to its previous behaviour otherwise, and every +write is best-effort -- a read-only home directory must never stop a run. + +Set ``WEIGHTSLAB_STATE_DIR`` to relocate the file (tests use it to stay out of +the developer's real state). +""" + +from __future__ import annotations + +import json +import logging +import os +import tempfile +from datetime import datetime, timezone +from pathlib import Path +from typing import Optional + +logger = logging.getLogger(__name__) + +_FILE_NAME = "active_experiment.json" +_SECTIONS = ("ui", "backend") +# A long-lived machine should not accumulate entries nobody can attribute. +_MAX_ENTRIES = 16 + + +def state_dir() -> Path: + """Directory holding the marker: ``$WEIGHTSLAB_STATE_DIR`` or ``~/.weightslab``.""" + override = (os.environ.get("WEIGHTSLAB_STATE_DIR") or "").strip() + if override: + return Path(override).expanduser() + return Path.home() / ".weightslab" + + +def state_path() -> Path: + """Absolute path of the marker file (may not exist yet).""" + return state_dir() / _FILE_NAME + + +def read_state() -> dict: + """The whole marker, or ``{}`` when it is absent or unreadable.""" + path = state_path() + try: + with open(path, "r", encoding="utf-8") as fh: + data = json.load(fh) + except FileNotFoundError: + return {} + except Exception as exc: # noqa: BLE001 -- a corrupt marker is not fatal + logger.debug("[active-experiment] could not read %s: %s", path, exc) + return {} + return data if isinstance(data, dict) else {} + + +def entries(section: str, state: Optional[dict] = None) -> list: + """One section as a list, accepting the older single-object shape.""" + value = (read_state() if state is None else state).get(section) + if isinstance(value, list): + return [e for e in value if isinstance(e, dict)] + if isinstance(value, dict): + return [value] + return [] + + +def _pid_is_running(pid) -> bool: + try: + pid = int(pid) + except (TypeError, ValueError): + return False + if pid <= 0: + return False + try: + import psutil + return psutil.pid_exists(pid) + except Exception: # noqa: BLE001 -- psutil missing/unusable: assume gone + return False + + +class _MarkerLock: + """Brief exclusive hold on the marker, via an O_EXCL lock file. + + Read-modify-write on a shared file loses updates when two processes do it + at once, and two processes doing it at once is precisely the case this + marker exists for: starting a classification UI and a segmentation UI + together dropped one of their port stamps, which is the field that tells + them apart afterwards. + + Best-effort by design: if the lock cannot be taken (a stale file nobody + cleaned up, a read-only home), the write proceeds anyway -- an advisory + record must never block a run. A stale lock older than a few seconds is + broken on purpose, since nothing here holds it for more than a file write. + """ + + STALE_AFTER = 5.0 + + def __init__(self, path: Path, attempts: int = 60, delay: float = 0.02): + self._path = Path(str(path) + ".lock") + self._attempts = attempts + self._delay = delay + self._held = False + + def __enter__(self): + import time + for _ in range(self._attempts): + try: + fd = os.open(str(self._path), os.O_CREAT | os.O_EXCL | os.O_WRONLY) + os.write(fd, str(os.getpid()).encode()) + os.close(fd) + self._held = True + return self + except FileExistsError: + try: + age = time.time() - os.path.getmtime(self._path) + if age > self.STALE_AFTER: + os.unlink(self._path) + continue + except OSError: + pass + time.sleep(self._delay) + except OSError: + break # cannot lock here at all; proceed unlocked + return self + + def __exit__(self, *exc): + if self._held: + try: + os.unlink(self._path) + except OSError: + pass + return False + + +def _write_section(section: str, root_log_dir, **meta) -> Optional[Path]: + """Record this process's entry in one section, keeping the others. + + Replaces the entry for this pid (a process re-recording, e.g. once its port + is known), drops entries whose process is gone, and leaves every other live + entry in place -- that is what lets two experiments run side by side. Held + under _MarkerLock so two processes starting together cannot lose each + other's entry. + """ + if section not in _SECTIONS: + raise ValueError(f"unknown section {section!r}") + if not root_log_dir: + return None + + pid = os.getpid() + entry = { + "root_log_dir": str(Path(root_log_dir).expanduser().resolve()), + "pid": pid, + "updated_at": datetime.now(timezone.utc).isoformat(timespec="seconds"), + } + entry.update({k: v for k, v in meta.items() if v is not None}) + + path = state_path() + try: + path.parent.mkdir(parents=True, exist_ok=True) + with _MarkerLock(path): + # Re-read INSIDE the lock: another process may have added its own + # entry since this function started. + state = read_state() + kept = [e for e in entries(section, state) + if e.get("pid") != pid and _pid_is_running(e.get("pid"))] + state[section] = (kept + [entry])[-_MAX_ENTRIES:] + # Written via a temp file in the same directory, then replaced, so a + # concurrent reader never sees a half-written document. + fd, tmp_name = tempfile.mkstemp(dir=str(path.parent), prefix=".active-", suffix=".json") + try: + with os.fdopen(fd, "w", encoding="utf-8") as fh: + json.dump(state, fh, indent=2) + os.replace(tmp_name, path) + except Exception: + try: + os.unlink(tmp_name) + except OSError: + pass + raise + except Exception as exc: # noqa: BLE001 -- advisory record, never fatal + logger.debug("[active-experiment] could not record %s dir in %s: %s", + section, path, exc) + return None + logger.debug("[active-experiment] recorded %s root_log_dir=%s", section, entry["root_log_dir"]) + return path + + +def record_ui_experiment(root_log_dir, ui_port: Optional[int] = None, + backend_port: Optional[int] = None) -> Optional[Path]: + """Record the directory ``weightslab start`` just established. + + ``backend_port`` is what makes two concurrent UIs distinguishable: it is + the gRPC port this UI proxies to, so it can later ask for the experiment + directory of ITS backend rather than of whichever one started last. + """ + return _write_section("ui", root_log_dir, ui_port=ui_port, backend_port=backend_port) + + +def record_backend_experiment(root_log_dir, grpc_port: Optional[int] = None) -> Optional[Path]: + """Record the directory training actually resolved (``wl.serve()``).""" + return _write_section("backend", root_log_dir, grpc_port=grpc_port) + + +def _entry_dir(section: str, entry: dict) -> Optional[str]: + value = (entry or {}).get("root_log_dir") + if not isinstance(value, str) or not value: + return None + if not os.path.isdir(value): + # A run whose directory was deleted (or a marker copied between + # machines) must not redirect anything. + logger.debug("[active-experiment] %s dir %s no longer exists; ignoring", section, value) + return None + return value + + +def _live_entries(section: str) -> list: + return [e for e in entries(section) + if _pid_is_running(e.get("pid")) and _entry_dir(section, e)] + + +def _sole_live_dir(section: str, port_key: str = "", port: Optional[int] = None) -> Optional[str]: + """The directory of the ONE live entry that matches, or None. + + With a port, an entry naming it wins outright -- that is how a UI finds ITS + backend rather than whichever backend started last. Otherwise a single live + entry is unambiguous and is used; two or more are not guessed between. + """ + live = _live_entries(section) + if port is not None and port_key: + matching = [e for e in live if e.get(port_key) == port] + if matching: + # The same port twice can only be stale bookkeeping: newest wins. + return _entry_dir(section, matching[-1]) + if len(live) == 1: + return _entry_dir(section, live[0]) + if len(live) > 1: + logger.info( + "[active-experiment] %d live %s experiments recorded (%s); not " + "guessing between them -- name the directory explicitly " + "(WEIGHTSLAB_ROOT_LOG_DIR, or root_log_dir in the config).", + len(live), section, + ", ".join(str(e.get("root_log_dir")) for e in live)) + return None + + +def ui_experiment_dir() -> Optional[str]: + """Directory of the most recent recorded ``weightslab start``, live or not. + + Raw record: it outlives the process that wrote it. Callers that REDIRECT a + run on this should use :func:`live_ui_experiment_dir` instead. + """ + for entry in reversed(entries("ui")): + found = _entry_dir("ui", entry) + if found: + return found + return None + + +def live_ui_experiment_dir() -> Optional[str]: + """Directory of a ``weightslab start`` that is *still running*. + + The handoff exists for "the UI is up over there, put this run in its + experiment". A record left behind by a UI that has since exited must not + silently redirect an unrelated run months later -- which is exactly what + happened to this repo's own gRPC tests: they resolved into a previous + session's experiment directory and loaded ITS config. And with two UIs up + (two experiments side by side) there is no right answer to guess. + """ + return _sole_live_dir("ui") + + +def backend_experiment_dir() -> Optional[str]: + """Directory of the most recent recorded ``wl.serve()``, live or not.""" + for entry in reversed(entries("backend")): + found = _entry_dir("backend", entry) + if found: + return found + return None + + +def live_backend_experiment_dir(grpc_port: Optional[int] = None) -> Optional[str]: + """Directory of a backend that is *still running*. + + Pass ``grpc_port`` -- the port the caller actually talks to -- and the + backend serving it is picked out by name. Without it, one live backend is + unambiguous and two are not guessed between: a UI showing the OTHER + experiment's reports is worse than a UI showing its own directory. + + Dead entries are ignored: the record outlives the process that wrote it. + """ + return _sole_live_dir("backend", "grpc_port", grpc_port) + + +def clear() -> None: + """Remove the marker (best-effort). Used by tests and ``weightslab`` teardown.""" + try: + state_path().unlink() + except FileNotFoundError: + pass + except Exception as exc: # noqa: BLE001 + logger.debug("[active-experiment] could not clear marker: %s", exc) diff --git a/weightslab/utils/logs.py b/weightslab/utils/logs.py index dd75872f..717165b4 100644 --- a/weightslab/utils/logs.py +++ b/weightslab/utils/logs.py @@ -9,7 +9,7 @@ # Define the log format to include timestamp, level, module name, and function name -FORMAT = '%(asctime)s.%(msecs)03d %(levelname)s:%(name)s:%(funcName)s: %(message)s' +FORMAT = '%(asctime)s.%(msecs)03d %(levelname)s:%(name)s:%(filename)s:%(lineno)d:%(funcName)s: %(message)s' DATE_FORMAT = '%d/%m/%Y-%H:%M:%S' # Global variables to track the log file path and handler From 0ab50a94bc1cf905999df1c7e1ee0dd71d7fb532 Mon Sep 17 00:00:00 2001 From: GuillaumePELLUET Date: Thu, 10 Sep 2026 15:24:38 +0200 Subject: [PATCH 10/29] test: fix the four failures in the unit suite Both pairs were test bugs, not defects in the code under test. tests/test_opencode_binary.py -- resolve_opencode_argv() returns str(Path), so the expected value has to be spelled the same way: str(Path("/mgd/opencode")) is "/mgd/opencode" on POSIX and "\mgd\opencode" on Windows. The POSIX form was hard-coded, so `test_managed_present_wins` and `test_download_when_no_path` failed on Windows only. Compare against str(Path(...)). tests/general/test_four_way_standalone.py -- `set_hp` refuses to guess when the process holds more than one hyperparam set ("Multiple hyperparam sets present; provide hp_name explicitly"), and the ledger is global: the sets registered by the other levels in this file, and by any module that ran earlier in the same process, are still there. So the two set_hp tests passed or failed depending on what had run before them. They now name their set via resolve_hp_name(), exactly as test_hp_lists_and_shows in the same file already does for `hp` -- and which its own comment explains was added for this reason. Full suite as CI runs it (pytest ./tests -m "not scale"): 1901 passed, 145 skipped, 17 deselected. Co-Authored-By: Claude Opus 5 (1M context) --- tests/general/test_four_way_standalone.py | 11 +++++++++-- tests/test_opencode_binary.py | 10 ++++++++-- 2 files changed, 17 insertions(+), 4 deletions(-) diff --git a/tests/general/test_four_way_standalone.py b/tests/general/test_four_way_standalone.py index 48a80e7d..f6ad5e3b 100644 --- a/tests/general/test_four_way_standalone.py +++ b/tests/general/test_four_way_standalone.py @@ -504,15 +504,22 @@ def test_hp_lists_and_shows(self): self.assertEqual(shown["name"], name) self.assertEqual(shown["hyperparams"]["experiment_name"], "standalone_config") + # `set_hp` refuses to guess when the process holds more than one + # hyperparam set ("Multiple hyperparam sets present; provide hp_name + # explicitly"), and the ledger is global: the sets registered by the other + # levels in this file (and by any test module that ran earlier in the same + # process) are still there. Name the set, exactly as test_hp_lists_and_shows + # already does for `hp` -- the alternative, asserting on whichever set the + # CLI happens to pick, is what made these two order-dependent. def test_set_hp_updates_the_live_config(self): - answer = self.cli("set_hp optimizer.lr 0.0005") + answer = self.cli(f"set_hp {resolve_hp_name()} optimizer.lr 0.0005") self.assertTrue(answer["ok"], answer) self.assertEqual(answer["key"], "optimizer.lr") self.assertEqual(answer["value"], 0.0005) self.assertEqual(self.hp["optimizer"]["lr"], 0.0005) def test_set_hp_updates_a_nested_data_key(self): - answer = self.cli("set_hp data.train_loader.batch_size 32") + answer = self.cli(f"set_hp {resolve_hp_name()} data.train_loader.batch_size 32") self.assertTrue(answer["ok"], answer) self.assertEqual(self.hp["data"]["train_loader"]["batch_size"], 32) diff --git a/tests/test_opencode_binary.py b/tests/test_opencode_binary.py index 18f48f3d..ed0ed1eb 100644 --- a/tests/test_opencode_binary.py +++ b/tests/test_opencode_binary.py @@ -218,11 +218,17 @@ class ResolverPrecedenceTests(unittest.TestCase): """opencode_process.resolve_opencode_argv order: managed-present -> PATH -> managed-download -> npx -> None.""" + # resolve_opencode_argv returns str(Path), so the expected value has to be + # spelled the same way: str(Path("/mgd/opencode")) is "/mgd/opencode" on + # POSIX and "\\mgd\\opencode" on Windows. Hard-coding the POSIX form made + # these two fail on Windows only, for no reason in the code under test. + MANAGED = str(Path("/mgd/opencode")) + def test_managed_present_wins(self): with patch.object(opencode_process.opencode_binary, "find_managed_binary", return_value=Path("/mgd/opencode")), \ patch.object(opencode_process.shutil, "which", return_value="/usr/bin/opencode"): - self.assertEqual(opencode_process.resolve_opencode_argv(), ["/mgd/opencode"]) + self.assertEqual(opencode_process.resolve_opencode_argv(), [self.MANAGED]) def test_path_used_before_download(self): with patch.object(opencode_process.opencode_binary, "find_managed_binary", return_value=None), \ @@ -237,7 +243,7 @@ def test_download_when_no_path(self): patch.object(opencode_process.opencode_binary, "ensure_managed_binary", return_value=Path("/mgd/opencode")), \ patch.object(opencode_process.shutil, "which", return_value=None): - self.assertEqual(opencode_process.resolve_opencode_argv(), ["/mgd/opencode"]) + self.assertEqual(opencode_process.resolve_opencode_argv(), [self.MANAGED]) def test_npx_last_resort(self): def which(name): From 5cd4920468a6f99b3b77617b27144b58c67aceb8 Mon Sep 17 00:00:00 2001 From: GuillaumePELLUET Date: Fri, 11 Sep 2026 12:17:50 +0200 Subject: [PATCH 11/29] Fix Ultralytics issue with v8.* --- weightslab/data/dataframe_manager.py | 121 +++++++++++++++++- .../integrations/ultralytics/trainer.py | 19 ++- weightslab/trainer/services/data_service.py | 28 ++++ 3 files changed, 162 insertions(+), 6 deletions(-) diff --git a/weightslab/data/dataframe_manager.py b/weightslab/data/dataframe_manager.py index 59afaf2c..064a6a7c 100644 --- a/weightslab/data/dataframe_manager.py +++ b/weightslab/data/dataframe_manager.py @@ -31,6 +31,84 @@ logger = logging.getLogger(__name__) # Set up logger +def label_is_empty(value) -> bool: + """True when a label cell carries nothing: None, NaN, or an empty sequence. + + Kept deliberately narrow — a scalar, a populated list and a populated array + are all "present", so a merge never overwrites a label a writer put there. + """ + if value is None: + return True + if isinstance(value, float) and value != value: # NaN + return True + if isinstance(value, (list, tuple, np.ndarray, dict)): + return len(value) == 0 + return False + + +def merge_instance_labels(values, sample_ids, annotation_ids) -> Dict[Any, list]: + """Group per-annotation label values back into one list per sample. + + ``_expand_records_to_multi_index`` splits a multi-instance label (N boxes, + N masks) across annotation rows 1..N and leaves the SAMPLE row's own value + empty — "reconstructed by the UI collapse", as its convention comment says. + This is that reconstruction: ``{sample_id: [inst_1, inst_2, ...]}`` ordered + by annotation_id, so a sample-centric consumer gets the whole list of boxes + rather than the first one (or, since the sample row is empty, none at all). + + Instances carrying no value are skipped; a sample with no instance value at + all is absent from the result, so callers can treat the mapping as "samples + that have something to merge". + """ + if values is None or len(values) == 0: + return {} + annots = np.asarray(annotation_ids) + try: + order = np.argsort(annots.astype(np.int64), kind="stable") + except (TypeError, ValueError): + order = np.arange(len(values)) + + merged: Dict[Any, list] = {} + for i in order: + value = values[i] + if label_is_empty(value): + continue + merged.setdefault(sample_ids[i], []).append( + value.tolist() if isinstance(value, np.ndarray) else value + ) + return merged + + +def fill_missing_labels(frame: pd.DataFrame, column: str, merged: Dict[Any, list]) -> bool: + """Write ``merged`` lists into ``frame[column]`` wherever the cell is empty. + + Only empty cells are filled: a sample row that already holds a label (the + single-instance layout, or a per-sample aggregate written after expansion) + keeps it. Returns True when anything was written. Touches only the rows + named in ``merged`` — never the whole frame. + """ + if not merged or column not in frame.columns or frame.empty: + return False + try: + keys = [k for k in merged if k in frame.index] + if not keys: + return False + positions = frame.index.get_indexer(pd.Index(keys)) + cells = frame[column].to_numpy(dtype=object, copy=True) + wrote = False + for key, pos in zip(keys, positions): + if pos < 0 or not label_is_empty(cells[pos]): + continue + cells[pos] = merged[key] + wrote = True + if wrote: + frame[column] = cells + return wrote + except Exception as exc: # noqa: BLE001 -- a merge must never fail a refresh + logger.debug("[merge_instance_labels] fill skipped for %r: %s", column, exc) + return False + + def _safe_update(target: pd.DataFrame, source: pd.DataFrame) -> None: """In-place update of ``target`` from ``source``, immune to the pandas internal AssertionError that ``DataFrame.update()`` raises when the source @@ -1490,7 +1568,15 @@ def clear_view_dirty(self): self._view_pending.clear() def get_source_rows(self, sample_ids, columns=None): - """Rows for *sample_ids* straight from the source frame. O(len(ids)).""" + """Rows for *sample_ids* straight from the source frame. O(len(ids)). + + ``columns`` is a request, not an assertion: a caller asking for a column + the source does not hold (a view-only column such as the natural-sort + ``signals.defaults.natural``) gets the columns that DO exist rather than + a KeyError. A column the source lacks has nothing to sync anyway, and + raising here failed the caller's whole request — GetDataSamples returned + `success=false` and the studio's modal came up empty. + """ with self._lock: if self._df.empty or not len(sample_ids): return None @@ -1499,7 +1585,16 @@ def get_source_rows(self, sample_ids, columns=None): want = set(str(s) for s in sample_ids) mask = keys.astype(str).isin(want) sub = self._df.loc[mask] - return sub[columns] if columns else sub + if not columns: + return sub + present = [c for c in columns if c in sub.columns] + if len(present) != len(columns): + logger.debug( + "[LedgeredDataFrameManager] get_source_rows: ignoring %d column(s) " + "absent from the source frame: %s", + len(columns) - len(present), + [c for c in columns if c not in sub.columns]) + return sub[present] def get_origin_revision(self, origin: str) -> int: # No lock: a dict read is atomic under the GIL, and this is polled from @@ -2738,6 +2833,28 @@ def get_collapse_annotations_to_samples_df(self, df: pd.DataFrame | None = None) sub_sid = sid_arr[inst_mask] sub_annot = annot_int[inst_mask] + # Multi-instance labels (one box / mask per annotation row) are + # merged back into a single list-of-lists on the sample row. + # `_expand_records_to_multi_index` deliberately leaves the sample + # row's target EMPTY for these samples and puts each instance on + # rows 1..N, so without this the sample-centric view carries no + # label at all for every multi-instance sample -- the studio drew + # zero boxes on exactly the samples that have several. + sub_sid_labels = sub_sid.tolist() + sub_annot_labels = sub_annot.tolist() + for label_col in (SampleStatsEx.TARGET.value, + SampleStatsEx.PREDICTION.value): + if label_col not in df.columns or label_col not in base.columns: + continue + # Vectorized skip: a column no instance row writes (predictions + # today) costs an O(n) C-level scan here instead of a .tolist() + # plus a Python pass over every annotation row. + if not bool(sub[label_col].notna().any()): + continue + merged_labels = merge_instance_labels( + sub[label_col].tolist(), sub_sid_labels, sub_annot_labels) + fill_missing_labels(base, label_col, merged_labels) + # Per-instance signal columns: numeric columns that carry any value on the # instance rows (written by enqueue_instance_batch). Array/io/meta excluded. exclude = { diff --git a/weightslab/integrations/ultralytics/trainer.py b/weightslab/integrations/ultralytics/trainer.py index 2fac84cc..868cca16 100644 --- a/weightslab/integrations/ultralytics/trainer.py +++ b/weightslab/integrations/ultralytics/trainer.py @@ -49,7 +49,9 @@ # ─── per-task wiring config ───────────────────────────────────────────── -# `loss_items` order per task (drives both channel names and index mapping): +# `loss_items` per task — a dict keyed "_loss" on UL >= 8.4.15x, a +# flat tensor in this same order before that (drives channel names, and the +# index mapping on the tensor path): # detect : (box, cls, dfl) # segment: (box, seg, cls, dfl, sem) ← we ship the first four; sem is 0 for # models without a semantic head. @@ -163,9 +165,18 @@ def _on_train_batch_end(trainer): ch = state["channels"] li = getattr(trainer, "loss_items", None) if li is not None and ch: - for i, n in enumerate(train_loss_names): - if i < li.numel(): - ch[f"train/{n}"](li[i:i+1].detach()) + # UL >= 8.4.15x hands back a {name: 0-dim tensor} dict keyed + # "box_loss"/"cls_loss"/...; older versions a flat tensor in + # the same order. Take names when we have them, index when we + # do not. + if isinstance(li, dict): + values = [li.get(f"{n}_loss", li.get(n)) for n in train_loss_names] + else: + values = [li[i:i+1] if i < li.numel() else None + for i in range(len(train_loss_names))] + for n, v in zip(train_loss_names, values): + if v is not None: + ch[f"train/{n}"](v.detach().reshape(1)) wl.guard_training_context.__exit__(None, None, None) def _on_val_batch_start(validator): diff --git a/weightslab/trainer/services/data_service.py b/weightslab/trainer/services/data_service.py index a2b0b477..498b8c24 100755 --- a/weightslab/trainer/services/data_service.py +++ b/weightslab/trainer/services/data_service.py @@ -27,6 +27,8 @@ DictConfig = dict # type: ignore from weightslab.data.sample_stats import SampleStatsEx +from weightslab.data.dataframe_manager import ( + merge_instance_labels, fill_missing_labels) from weightslab.utils.tools import safe_reset_index from weightslab.data.h5_dataframe_store import H5DataFrameStore from weightslab.proto.experiment_service_pb2 import SampleEditType @@ -4092,6 +4094,7 @@ def _fastUpdateInternals(self, max_dirty: int = 250_000) -> bool: # Python, so GetDataSamples then reported discarded="1" and the studio # greyed the sample out, progressively, exactly as the model worked # through the dataset -- while the dataframe itself still said False. + merged_labels: dict = {} if isinstance(sub.index, pd.MultiIndex): ANNOT = SampleStatsEx.INSTANCE_ID.value names = list(getattr(sub.index, "names", []) or []) @@ -4101,6 +4104,27 @@ def _fastUpdateInternals(self, max_dirty: int = 250_000) -> bool: canonical = np.asarray(annot).astype(int) == 0 except (TypeError, ValueError): canonical = np.array([str(a) in ("0", "0.0") for a in annot]) + # Multi-instance labels: the annotation rows about to be dropped + # hold one box/mask each, and the canonical row's own label is + # EMPTY by construction (_expand_records_to_multi_index). Merge + # them exactly as get_collapse_annotations_to_samples_df does -- + # otherwise this differential write puts that empty value straight + # over the merged list the rebuild had placed in the view, and the + # sample loses its boxes again on the next refresh. + inst = ~canonical + if inst.any(): + inst_sids = sub.index.get_level_values(0).to_numpy()[inst].tolist() + inst_annot = np.asarray(annot).tolist() + inst_annot = [a for a, keep in zip(inst_annot, inst) if keep] + for label_col in (SampleStatsEx.TARGET.value, + SampleStatsEx.PREDICTION.value): + if label_col not in sub.columns: + continue + merged = merge_instance_labels( + sub[label_col].to_numpy(dtype=object)[inst].tolist(), + inst_sids, inst_annot) + if merged: + merged_labels[label_col] = merged if canonical.any(): sub = sub[canonical] sub = sub.droplevel(-1) @@ -4108,6 +4132,10 @@ def _fastUpdateInternals(self, max_dirty: int = 250_000) -> bool: # row when the level was present, the first occurrence otherwise) -- # never the last, for the reason above. sub = sub[~sub.index.duplicated(keep="first")] + if merged_labels: + sub = sub.copy() + for label_col, merged in merged_labels.items(): + fill_missing_labels(sub, label_col, merged) # Only rows the view actually holds; a structural change (new sample) # must still fall back to the full rebuild rather than be invented here. From a2b89029aabb446a0327127f6be7d5667411bb20 Mon Sep 17 00:00:00 2001 From: GuillaumePELLUET Date: Wed, 23 Sep 2026 01:13:49 +0200 Subject: [PATCH 12/29] test: fix the gRPC serve flake that failed CI with identical code test_grpc_serve_honors_explicit_port_without_force_parameters timed out on commit ac405334's PR run (35738307562) and passed on its push run (35738302340) -- same commit, same workflow. It was never deadlocked. serving_thread_callback logs its way through the bind, and in a full-suite run each of those calls takes SECONDS: 00:35:29.215 [gRPC] Thread callback started 00:35:39.294 [gRPC] Creating ThreadPoolExecutor <- +10.1s 00:35:44.893 [gRPC] Server object created <- +5.6s Three log lines, 15 seconds. Every assertion would have passed; the test just never reached them before _TimeoutMixin fired -- hence a timeout rather than a failure, and a green rerun every time. The contention comes from threads earlier tests start and never stop (the embedded ipykernel's tornado/zmq loop, dataframe_manager's flush threads and their DEBUG chatter): by the time pytest reaches tests/trainer/ they hold the logging lock often enough to starve it. tests/trainer/services on its own passes -- 386 passed. - stub the module logger in _run_grpc_serve_capturing_bind. These tests assert on add_insecure_port and have no interest in log output, so the fix is to stop depending on real logging rather than to outwait it. - _TimeoutMixin built its exc_info with a None traceback, so pytest raised "'NoneType' object is not iterable" while rendering and fell back to "Incompatible Exception Representation" -- the report named no location at all. Raise and catch instead, for a real traceback. - 30s -> 60s: the cap guards against a stuck thread, it shouldn't double as a performance assertion. This alone did NOT fix the hang. - conftest: keep the real ResourceMonitor out of the suite. grpc_serve ends by calling start_resource_monitor_from_config() and nothing stubbed it, so the first such test left a process-wide singleton sampling CPU/memory/disk/ network/GPU for the rest of the session. One of the leaked threads above. Full suite: 1847 passed, 145 skipped, 0 failed. Co-Authored-By: Claude Opus 5 (1M context) --- tests/conftest.py | 33 ++++++++++ .../services/test_trainer_services_server.py | 65 +++++++++++++++++-- 2 files changed, 93 insertions(+), 5 deletions(-) diff --git a/tests/conftest.py b/tests/conftest.py index 32c86157..30bcc6f0 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -35,3 +35,36 @@ def _isolate_weightslab_state(): os.environ.pop("WEIGHTSLAB_STATE_DIR", None) else: os.environ["WEIGHTSLAB_STATE_DIR"] = previous + + +@pytest.fixture(scope="session", autouse=True) +def _disable_resource_monitor(): + """Keep the real resource monitor out of the suite. + + ``grpc_serve`` ends by calling ``start_resource_monitor_from_config()``, and + none of the tests that exercise it stub that out -- so the first such test + starts a REAL ResourceMonitor: a process-wide singleton whose sampling + thread then polls CPU/memory/disk/network/GPU for the rest of the session, + logging as it goes. Nothing stops it, because the singleton makes every + later call a no-op that returns the already-running instance. + + That background load is not free on a 2-core CI runner. It showed up as + tests/trainer/services/test_trainer_services_server.py's + test_grpc_serve_honors_explicit_port_without_force_parameters timing out + against its 30s cap -- while the SAME commit passed on the push run + (ac405334: PR run 35738307562 failed, push run 35738302340 passed), which + is the signature of contention rather than a real defect. + + The monitor's own tests are unaffected: they patch + ``load_resource_monitoring_config`` directly, or set this same variable + inside their own ``_patched_env``. + """ + previous = os.environ.get("WEIGHTSLAB_DISABLE_RESOURCE_MONITORING") + os.environ["WEIGHTSLAB_DISABLE_RESOURCE_MONITORING"] = "1" + try: + yield + finally: + if previous is None: + os.environ.pop("WEIGHTSLAB_DISABLE_RESOURCE_MONITORING", None) + else: + os.environ["WEIGHTSLAB_DISABLE_RESOURCE_MONITORING"] = previous diff --git a/tests/trainer/services/test_trainer_services_server.py b/tests/trainer/services/test_trainer_services_server.py index 851478a5..9b8cc41d 100644 --- a/tests/trainer/services/test_trainer_services_server.py +++ b/tests/trainer/services/test_trainer_services_server.py @@ -1,3 +1,4 @@ +import sys import unittest from concurrent.futures import ThreadPoolExecutor, TimeoutError as FuturesTimeoutError from unittest.mock import MagicMock, patch @@ -6,7 +7,15 @@ # Default per-test timeout in seconds. Override with WL_TEST_TIMEOUT env var. import os -_TEST_TIMEOUT = int(os.getenv("WL_TEST_TIMEOUT", "30")) + +# 60s, not 30s. These tests are pure-mock and finish in milliseconds locally, +# so the cap exists only to stop a genuinely stuck gRPC thread from hanging CI +# -- it was never meant to be a performance assertion. At 30s it became one: +# commit ac405334 timed out here on its PR run (35738307562) and passed on the +# push run (35738302340) with identical code, purely on runner contention. +# Doubling keeps the runaway-thread protection while leaving room for a loaded +# 2-core runner. pytest's own --timeout=600 is still the outer backstop. +_TEST_TIMEOUT = int(os.getenv("WL_TEST_TIMEOUT", "60")) class _TimeoutMixin: @@ -20,11 +29,27 @@ def run(self, result=None): fut.result(timeout=_TEST_TIMEOUT) except FuturesTimeoutError: if result is not None: - result.addError(self, (TimeoutError, TimeoutError( - f"Test timed out after {_TEST_TIMEOUT}s"), None)) + result.addError(self, self._timeout_exc_info()) finally: pool.shutdown(wait=False) + @staticmethod + def _timeout_exc_info(): + """Build a real (type, value, traceback) triple for the timeout. + + Passing ``None`` as the traceback -- which this did -- is what turned a + timeout into an unreadable CI failure: pytest tries to walk the + traceback to render the entry, raises "'NoneType' object is not + iterable" doing it, and falls back to "NOTE: Incompatible Exception + Representation", so the report says only that something timed out and + nothing about where. Raising and catching gives us a genuine traceback + object, so the failure renders normally. + """ + try: + raise TimeoutError(f"Test timed out after {_TEST_TIMEOUT}s") + except TimeoutError: + return sys.exc_info() + class TestExperimentServiceServicerDelegation(_TimeoutMixin, unittest.TestCase): def test_servicer_delegates_to_subservices(self): @@ -304,7 +329,36 @@ def start(self): def _run_grpc_serve_capturing_bind(self, bound_port, **serve_kwargs): """Run grpc_serve with the same scaffolding as the tests above and return - the fake server, so a caller can assert which address was bound.""" + the fake server, so a caller can assert which address was bound. + + The module logger is stubbed along with everything else, and that stub + is what makes these two tests survive a full-suite run. They assert on + ``add_insecure_port`` -- they have no interest in log output -- but + ``serving_thread_callback`` logs its way through the bind, and by the + time pytest reaches this file (late: ``tests/trainer/`` sorts after + backend, components, data, export, gRPC, general, integrations, model, + modules, monitoring) the process is full of background threads that + earlier tests started and never stopped -- the embedded ipykernel's + tornado/zmq loop, dataframe_manager's flush threads and their constant + DEBUG chatter. They hold the logging lock often enough that a single + ``logger.info`` starts taking SECONDS: + + 00:35:29.215 [gRPC] Thread callback started + 00:35:39.294 [gRPC] Creating ThreadPoolExecutor <- +10.1s + 00:35:44.893 [gRPC] Server object created <- +5.6s + + Three log lines, 15 seconds. Nothing is deadlocked and the assertions + would all pass -- the test just never gets to them before _TimeoutMixin + fires, which is why it dies with a timeout rather than a failure, and + why it passes in isolation and on a rerun. That is the CI flake: + commit ac405334 failed this on its PR run (35738307562) and passed on + its push run (35738302340). + + Raising the timeout only buys time against an unbounded queue; removing + the dependency on real logging fixes it. The leaked threads themselves + are a separate (real) problem -- see tests/conftest.py's resource + monitor note for one of them. + """ fake_server = MagicMock() fake_server.add_insecure_port.return_value = bound_port @@ -349,7 +403,8 @@ def start(self): patch("weightslab.trainer.trainer_services.WeighlabsWatchdog", return_value=fake_watchdog), \ patch("weightslab.trainer.trainer_services.time.sleep", side_effect=_sleep_stop), \ patch("weightslab.trainer.trainer_services.ExperimentServiceServicer"), \ - patch("weightslab.trainer.trainer_services.get_hyperparams", return_value={}): + patch("weightslab.trainer.trainer_services.get_hyperparams", return_value={}), \ + patch("weightslab.trainer.trainer_services.logger"): trainer_services.grpc_serve(**serve_kwargs) return fake_server From caa1d7478c56c8a11d71cdeb0d4aeea04a971029 Mon Sep 17 00:00:00 2001 From: GuillaumePELLUET Date: Wed, 23 Sep 2026 10:45:05 +0200 Subject: [PATCH 13/29] fix(notebook): shut the embedded kernel down instead of aborting at exit CI died on the way out of a run whose tests had all passed: Fatal Python error: _enter_buffered_busy: could not acquire lock for <_io.BufferedWriter name=''> at interpreter shutdown, possibly due to daemon threads ... Aborted (core dumped) The embedded Jupyter kernel runs as a daemon thread on app.start() -- a tornado/zmq loop that never returns -- and had no shutdown path at all. So it was still live when CPython began finalizing, and the shutdown raced itself: the interpreter closes the zmq sockets, tornado's zmqstream notices ("Got events for stream ... attached to closed socket") and calls gen_log.warning, and logging writes into a stderr buffer whose lock is already being torn down. That aborts the process; the job then fails on the exit code even though the suite passed. Not test-only, which is why this is here and not in tests/conftest.py: any script that enables the notebook takes the same risk at exit. Stop the kernel from an atexit hook, which runs BEFORE finalization, so app.start() returns and the thread exits while logging still works. The loop to stop is not the obvious one. _run_embedded_kernel sets up an asyncio loop for its thread, but ipykernel builds its own AsyncIOMainLoop under app.io_loop and blocks on that instead -- stopping ours is a no-op (measured: our loop ...692880 vs app's ...181904, thread still alive after a 10s join). Go through app.io_loop, whose add_callback is thread-safe, and join with a bounded timeout so a stubborn kernel can't hang the process on the way out. Verified: thread alive -> not alive across the hook; repro exits 0 with zero "Fatal Python error". Note the abort does not reproduce on Windows -- it is a platform- and timing-dependent finalization race -- so what is proven locally is the mechanism, not the end-to-end CI symptom. Full suite unchanged: 1847 passed, 145 skipped, 0 failed. Co-Authored-By: Claude Opus 5 (1M context) --- .../trainer/services/notebook_service.py | 69 +++++++++++++++++-- 1 file changed, 65 insertions(+), 4 deletions(-) diff --git a/weightslab/trainer/services/notebook_service.py b/weightslab/trainer/services/notebook_service.py index 5679d741..c3143160 100644 --- a/weightslab/trainer/services/notebook_service.py +++ b/weightslab/trainer/services/notebook_service.py @@ -31,6 +31,7 @@ import sys import json import time +import atexit import ctypes import queue import shutil @@ -300,13 +301,63 @@ def _safe(getter): # "done" -- keying off "started" makes them give up the instant the attempt is # launched, before the kernel has had any chance to write its connection file. _EMBED_STATE = {"started": False, "connection_file": None, "kernel_thread_id": None, - "done": threading.Event()} + "done": threading.Event(), "thread": None} # Rebound on every NotebookService construction (i.e. every watchdog restart) # so the one long-lived embedded kernel thread always refreshes df/root_log_dir # against whichever data_service/root_log_dir is currently live. _ACTIVE_BINDING = {"data_service": None, "root_log_dir": None} +def _shutdown_embedded_kernel(timeout: float = 5.0) -> None: + """Stop the embedded kernel's event loop and wait for its thread to finish. + + Registered with ``atexit`` when the kernel starts, because the thread is a + daemon running ``app.start()`` -- a tornado/zmq loop that never returns on + its own. Without this, the thread is still live when CPython begins + finalizing, and the shutdown sequence races itself: the interpreter closes + the zmq sockets, tornado's zmqstream notices ("Got events for stream ... + attached to closed socket") and calls ``gen_log.warning``, and logging then + writes to a stderr whose buffer lock is already being torn down. That is a + hard abort, not an exception:: + + Fatal Python error: _enter_buffered_busy: could not acquire lock for + <_io.BufferedWriter name=''> at interpreter shutdown, possibly + due to daemon threads + ... Aborted (core dumped) + + It killed a CI run whose tests had all passed -- the suite finished, the + process died on the way out, and the job failed on the exit code. It is not + test-only: any script that enables the notebook takes the same risk at + exit, which is why this lives here and not in tests/conftest.py. + + atexit runs before interpreter finalization, so stopping the loop here lets + app.start() return and the thread exit while logging still works. + """ + thread = _EMBED_STATE.get("thread") + if thread is None or not thread.is_alive(): + return + + # Stop the loop app.start() is ACTUALLY blocked on, which is not the one + # this module creates: _run_embedded_kernel sets up an asyncio loop for the + # thread, but ipykernel builds its own AsyncIOMainLoop underneath + # app.io_loop, and that is what runs. Stopping the loop we made is a no-op + # -- measured: stored loop id ...692880 vs app's ...181904, and the thread + # stayed alive through a 10s join. IOLoop.add_callback is documented + # thread-safe, which is what lets us reach into it from here. + try: + from ipykernel.kernelapp import IPKernelApp + io_loop = getattr(IPKernelApp.instance(), "io_loop", None) + if io_loop is not None: + io_loop.add_callback(io_loop.stop) + except Exception: + # ipykernel missing, never initialized, or already torn down. + pass + + # Bounded: a kernel that refuses to stop must not hang the process on the + # way out. Worst case we are back to the old behaviour. + thread.join(timeout=timeout) + + def configure_embedded_kernel(enabled: bool) -> None: """Record wl.serve()'s enable/disable decision for the embedded kernel.""" global _EMBED_ENABLED @@ -371,10 +422,17 @@ def ensure_embedded_kernel(data_service, root_log_dir: Path) -> None: _EMBED_STATE["done"].set() return connection_file = _ACTIVE_BINDING["root_log_dir"] / "notebook_kernel.json" - threading.Thread( + _kernel_thread = threading.Thread( target=_run_embedded_kernel, args=(connection_file,), name="WL-Embedded-Jupyter-Kernel", daemon=True, - ).start() + ) + _EMBED_STATE["thread"] = _kernel_thread + # Registered only once the kernel actually starts, and only here -- + # this is the one place that creates the thread. See + # _shutdown_embedded_kernel() for why a daemon thread left running into + # interpreter finalization aborts the process. + atexit.register(_shutdown_embedded_kernel) + _kernel_thread.start() try: if _wait_for_connection_file(connection_file, timeout=15.0): _EMBED_STATE["connection_file"] = connection_file @@ -500,7 +558,10 @@ def _run_embedded_kernel(connection_file: Path) -> None: # execute_interactive() and blocks waiting for it). _EMBED_STATE["kernel_thread_id"] = threading.get_ident() - # ipykernel's IOLoop needs a running asyncio loop on THIS thread. + # ipykernel's IOLoop needs a running asyncio loop on THIS thread. Note this + # is NOT the loop app.start() ends up blocked on -- ipykernel builds its own + # AsyncIOMainLoop under app.io_loop -- so _shutdown_embedded_kernel() goes + # through the app, not through this. asyncio.set_event_loop(asyncio.new_event_loop()) ns = build_notebook_namespace( From e9f7f238efb7227ffbaa4724952515fc698a21cb Mon Sep 17 00:00:00 2001 From: GuillaumePELLUET Date: Sat, 26 Sep 2026 00:56:59 +0200 Subject: [PATCH 14/29] fix utest error --- tests/components/test_checkpoint_workflow.py | 126 ++++++++++++++++++ .../services/test_trainer_services_unit.py | 2 + 2 files changed, 128 insertions(+) diff --git a/tests/components/test_checkpoint_workflow.py b/tests/components/test_checkpoint_workflow.py index 72bf7a90..c12d094e 100644 --- a/tests/components/test_checkpoint_workflow.py +++ b/tests/components/test_checkpoint_workflow.py @@ -1552,6 +1552,132 @@ def test_multiroot_adopts_most_recently_updated_nested_root(self): finally: lg.stop_background_flush() + def _write_viewer_root(self, root_dir, timestamp): + """An empty root with the newest manifest -- e.g. a read-only viewer's own + folder next to finished experiment roots -- so the manager adopts it as + the effective root and every experiment lives in a sibling root.""" + checkpoints_dir = root_dir / "checkpoints" + checkpoints_dir.mkdir(parents=True, exist_ok=True) + with open(checkpoints_dir / "manifest.yaml", 'w') as f: + yaml.dump({'experiments': {}, 'latest_hash': None, 'last_updated': timestamp}, f) + + def test_load_checkpoint_resolves_hash_from_sibling_root(self): + """A sibling root only contributes its curves to the effective root, but + its checkpoints stay where they are: loading one of its hashes must read + its weights from that sibling (latest, and closest to a target step).""" + parent = self.tmp_dir / "compare_exp" + root_a = parent / "goldset" + root_b = parent / "signal" + hash_a = "aaaa0010" + "bbbb0010" + "cccc0010" + hash_b = "aaaa0011" + "bbbb0011" + "cccc0011" + self._write_experiment_root(root_a, hash_a, step=20, timestamp="2024-01-01T00:00:00", loss_values=[0.3]) + extra = CheckpointManager(root_log_dir=str(root_a), load_model=False, load_config=False, load_data=False) + extra.current_exp_hash = hash_a + extra.save_model_checkpoint(model=nn.Sequential(nn.Linear(2, 2)), step=60, save_optimizer=False) + ledgers.clear_all() + self._pin_manifest_timestamp(root_a, hash_a, "2024-01-01T00:00:00") + self._write_experiment_root(root_b, hash_b, step=80, timestamp="2024-02-01T00:00:00", loss_values=[0.4]) + + manager = CheckpointManager(root_log_dir=str(parent), load_model=False, load_config=False, load_data=False) + self.assertEqual(manager.root_log_dir, root_b.absolute(), "the most recent root is the effective one") + + latest = manager.load_checkpoint(hash_a, load_model=False, load_config=False, load_data=False) + self.assertIn('weights', latest['loaded_components'], "the sibling root's weights should be found") + self.assertEqual(latest['weights'].get('step'), 60) + + at_step = manager.load_checkpoint(hash_a, load_model=False, load_config=False, load_data=False, target_step=25) + self.assertEqual(at_step['weights'].get('step'), 20, "closest checkpoint to the target step, in the sibling") + + own = manager.load_checkpoint(hash_b, load_model=False, load_config=False, load_data=False) + self.assertEqual(own['weights'].get('step'), 80, "the effective root's own hashes are unaffected") + + def test_load_state_restores_sibling_hash_from_an_empty_viewer_root(self): + """The reported case: the effective root is an empty viewer folder and the + experiment lives in a sibling root. Restoring its hash must apply the + weights and sync the current hash and component hashes with it.""" + parent = self.tmp_dir / "viewer_exp" + root_a = parent / "goldset_6pct" + hash_a = "aaaa0012" + "bbbb0012" + "cccc0012" + self._write_experiment_root(root_a, hash_a, step=6000, timestamp="2024-01-01T00:00:00", loss_values=[0.2]) + self._write_viewer_root(parent / "_studio", timestamp="2025-01-01T00:00:00") + + manager = CheckpointManager(root_log_dir=str(parent), load_model=False, load_config=False, load_data=False) + self.assertEqual(manager.root_log_dir, (parent / "_studio").absolute()) + self.assertTrue(manager.is_multi_root) + + model = nn.Sequential(nn.Linear(2, 2)) + ledgers.register_model(model) + ledgers.register_checkpoint_manager(manager) + lg = LoggerQueue(register=True) # a live session always has one; load_state reads its length + try: + ok = manager.load_state(hash_a, load_model=False, load_config=False, load_data=False, load_logger=False) + finally: + lg.stop_background_flush() + self.assertTrue(ok, "a hash living in a sibling root should be restorable") + self.assertEqual(manager.current_exp_hash, hash_a) + + saved = th.load(root_a / "checkpoints" / "models" / "bbbb0012" / f"{hash_a}_step_006000.pt", weights_only=False) + for name, tensor in saved['model_state_dict'].items(): + self.assertTrue(th.equal(model.state_dict()[name], tensor), f"weights of {name} should be restored") + + components = manager.hash_generator.get_component_hashes() + self.assertEqual((components['hp'], components['model'], components['data']), + ("aaaa0012", "bbbb0012", "cccc0012")) + + def test_load_state_brings_the_sibling_runs_per_sample_stats(self): + """The data snapshot only holds sample ids, tags and discards; a run's + per-sample stats (last loss, prediction, nb_seen...) live in its root's + data.h5. Restoring a sibling's hash must load them into the dataframe.""" + from weightslab.data.dataframe_manager import LedgeredDataFrameManager + from weightslab.data.h5_dataframe_store import H5DataFrameStore + + parent = self.tmp_dir / "viewer_df_exp" + root_a = parent / "goldset_6pct" + hash_a = "aaaa0014" + "bbbb0014" + "cccc0014" + self._write_experiment_root(root_a, hash_a, step=6000, timestamp="2024-01-01T00:00:00", loss_values=[0.2]) + H5DataFrameStore(root_a / "checkpoints" / "data" / "data.h5").upsert("train_loader", pd.DataFrame({ + "sample_id": [0, 1, 2], + "annotation_id": [0, 0, 0], + "discarded": [False, True, False], + "nb_seen": [8, 0, 8], + "signals//train-loss-CE": [0.1, float("nan"), 2.5], + }).set_index(["sample_id", "annotation_id"])) + self._write_viewer_root(parent / "_studio", timestamp="2025-01-01T00:00:00") + + manager = CheckpointManager(root_log_dir=str(parent), load_model=False, load_config=False, load_data=False) + dfm = LedgeredDataFrameManager(enable_flushing_threads=False, enable_h5_persistence=False) + dfm.upsert_df(pd.DataFrame({"sample_id": [0, 1, 2], "annotation_id": [0, 0, 0], "origin": "train_loader", + "discarded": False}).set_index(["sample_id", "annotation_id"]), + origin="train_loader") + ledgers.register_dataframe(dfm) + ledgers.register_model(nn.Sequential(nn.Linear(2, 2))) + ledgers.register_checkpoint_manager(manager) + lg = LoggerQueue(register=True) + try: + ok = manager.load_state(hash_a, load_model=False, load_config=False, load_data=True, load_logger=False) + finally: + lg.stop_background_flush() + + self.assertTrue(ok) + df = dfm.get_df_view() + self.assertAlmostEqual(df.loc[("0", 0), "signals//train-loss-CE"], 0.1) + self.assertAlmostEqual(df.loc[("2", 0), "signals//train-loss-CE"], 2.5) + self.assertEqual(int(df.loc[("0", 0), "nb_seen"]), 8) + self.assertTrue(bool(df.loc[("1", 0), "discarded"])) + + def test_unknown_hash_is_still_reported_missing_in_multiroot_mode(self): + parent = self.tmp_dir / "unknown_hash_exp" + self._write_experiment_root(parent / "a", "aaaa0013" + "bbbb0013" + "cccc0013", step=5, + timestamp="2024-01-01T00:00:00", loss_values=[0.1]) + self._write_viewer_root(parent / "_studio", timestamp="2025-01-01T00:00:00") + manager = CheckpointManager(root_log_dir=str(parent), load_model=False, load_config=False, load_data=False) + + unknown = "ffff0000" + "eeee0000" + "dddd0000" + result = manager.load_checkpoint(unknown, load_model=False, load_data=False) + self.assertEqual(result['loaded_components'], set()) + self.assertFalse(manager.load_state(unknown, load_logger=False)) + self.assertNotEqual(manager.current_exp_hash, unknown) + def test_sibling_merge_is_idempotent_across_call_sites(self): """merge_from_disk can be triggered from more than one init-ordering call site (logger created before vs. after the checkpoint manager); diff --git a/tests/trainer/services/test_trainer_services_unit.py b/tests/trainer/services/test_trainer_services_unit.py index a7c553dd..a95fc920 100644 --- a/tests/trainer/services/test_trainer_services_unit.py +++ b/tests/trainer/services/test_trainer_services_unit.py @@ -283,6 +283,8 @@ def test_restore_checkpoint_weights_step_mode(self): checkpoint_manager.load_state.assert_called_once() _, kwargs = checkpoint_manager.load_state.call_args self.assertEqual(kwargs.get("target_step"), 5) + # The grid view is rebuilt so every restored per-sample column shows up + service.data_service._slowUpdateInternals.assert_called_once_with(force=True) def _make_save_service(self, components): ctx = _DummyCtx(components=components) From 9927abc4217082782e1213b71d1554000db2e501 Mon Sep 17 00:00:00 2001 From: GuillaumePELLUET Date: Sat, 26 Sep 2026 00:58:09 +0200 Subject: [PATCH 15/29] update cls example with basic CNN 1M parameters --- .../PyTorch/wl-classification/main.py | 698 +++++++++++++++--- 1 file changed, 604 insertions(+), 94 deletions(-) diff --git a/weightslab/examples/PyTorch/wl-classification/main.py b/weightslab/examples/PyTorch/wl-classification/main.py index 5741684e..519975d1 100644 --- a/weightslab/examples/PyTorch/wl-classification/main.py +++ b/weightslab/examples/PyTorch/wl-classification/main.py @@ -1,5 +1,7 @@ import itertools +import math import os +import random import ssl import time import logging @@ -14,8 +16,12 @@ import yaml import tqdm +import json +import numpy as np +import pandas as pd import torch import torch.nn as nn +import torch.nn.functional as F import torch.optim as optim from torchvision import datasets, transforms @@ -24,7 +30,6 @@ from torch.utils.data import Dataset import weightslab as wl -from weightslab.examples.utils.baseline_models.pytorch.models import FashionCNN as CNN # Setup logging @@ -119,17 +124,126 @@ def __getitem__(self, idx): return image, idx, label +# ============================================================================= +# Model and recipe +# ============================================================================= +class CNN(nn.Module): + """2 conv + 2 fc layers (1.2M parameters), raw logits out, optional dropout. + + Raw logits, not a softmax: `nn.CrossEntropyLoss` applies its own log-softmax. + A softmax here would be applied twice, which still trains but squashes the + per-sample loss values -- and those values are what you sort and filter on + in the Studio. + """ + + def __init__(self, dropout=False): + super().__init__() + self.input_shape = (1, 1, 28, 28) + self.conv1 = nn.Conv2d(1, 32, 3) + self.relu1 = nn.ReLU() + self.conv2 = nn.Conv2d(32, 64, 3) + self.relu2 = nn.ReLU() + self.pool = nn.MaxPool2d(2) + self.drop1 = nn.Dropout(0.25 if dropout else 0.0) + self.flatten = nn.Flatten() + self.fc1 = nn.Linear(64 * 12 * 12, 128) + self.relu3 = nn.ReLU() + self.drop2 = nn.Dropout(0.5 if dropout else 0.0) + self.fc2 = nn.Linear(128, 10) + + def features(self, x): + x = self.drop1(self.pool(self.relu2(self.conv2(self.relu1(self.conv1(x)))))) + return self.relu3(self.fc1(self.flatten(x))) + + def forward(self, x): + return self.fc2(self.drop2(self.features(x))) + + +def gpu_augment(x, rot_deg=10.0, shift=0.1, scale=(0.9, 1.1)): + """Random affine per image (rotation, shift, scale), on the batch, on the GPU.""" + b = x.size(0) + ang = (torch.rand(b, device=x.device) * 2 - 1) * rot_deg * math.pi / 180 + s = torch.empty(b, device=x.device).uniform_(*scale) + tx = (torch.rand(b, device=x.device) * 2 - 1) * shift * 2 + ty = (torch.rand(b, device=x.device) * 2 - 1) * shift * 2 + cos, sin = torch.cos(ang) / s, torch.sin(ang) / s + theta = torch.stack([torch.stack([cos, -sin, tx], 1), torch.stack([sin, cos, ty], 1)], 1) + grid = F.affine_grid(theta, x.shape, align_corners=False) + return F.grid_sample(x, grid, align_corners=False, padding_mode="zeros") + + +def lr_at(step, total_steps, base_lr, sched): + """const, or warmcos: linear warmup over 15% of the steps, then cosine decay to ~0.""" + if sched == "const": + return base_lr + warm = max(1, int(0.15 * total_steps)) + if step <= warm: + return base_lr * (0.04 + 0.96 * step / warm) + p = (step - warm) / max(1, total_steps - warm) + return 1e-6 + base_lr * 0.5 * (1 + math.cos(math.pi * p)) + + +def get_val_ids(targets, path): + """Fixed validation split: 500 train images per class, never trained on. + + Fixed RNG seed, so the split is identical on every machine and every rerun. + Delete the file to draw a new one. + """ + if path and os.path.exists(path): + return sorted(int(l) for l in open(path) if l.strip()) + y = targets.numpy() if hasattr(targets, "numpy") else np.asarray(targets) + rng = np.random.RandomState(2026) + ids = [] + for c in range(10): + ids += rng.choice(np.where(y == c)[0], 500, replace=False).tolist() + ids = sorted(int(i) for i in ids) + if path: + open(path, "w").write("\n".join(map(str, ids)) + "\n") + return ids + + +class SubsetWithIds(Dataset): + """A view over `base` restricted to `ids`, in the order given. + + WeightsLab numbers each split's samples contiguously, so a sample id is a + position within the split, not an MNIST index. `ids` is the lookup back: + the n-th sample of this loader is `ids[n]` in MNIST, which is also the n-th + line of the matching train_ids.txt / val_ids.txt. + """ + + def __init__(self, base, ids): + self.base = base + self.ids = list(ids) + + def __len__(self): + return len(self.ids) + + def __getitem__(self, i): + return self.base[self.ids[i]] + + # ----------------------------------------------------------------------------- # Train / Test functions # ----------------------------------------------------------------------------- -def train(loader, model, optimizer, criterion_mlt, device): - """Single training step using the tracked dataloader + watched loss.""" +def train(loader, model, optimizer, criterion_mlt, device, step, recipe): + """Single training step using the tracked dataloader + watched loss. + + `recipe` carries the run's knobs: augmentation and the LR schedule. + """ with wl.guard_training_context: (inputs, ids, labels) = next(loader) inputs = inputs.to(device) labels = labels.to(device) + # Augment on the GPU, on the batch + if recipe["augment"]: + inputs = gpu_augment(inputs) + + # LR schedule is applied per step, from the model's age + for g in optimizer.param_groups: + g["lr"] = lr_at(step, recipe["total_steps"], recipe["lr"], recipe["schedule"]) + # Infer optimizer.zero_grad() preds_raw = model(inputs) @@ -140,7 +254,9 @@ def train(loader, model, optimizer, criterion_mlt, device): else: preds = preds_raw.argmax(dim=1, keepdim=True) - # Loss is a watched object => pass metadata for logging/stats + # Loss is a watched object => pass metadata for logging/stats. + # per_sample=True on the criterion keeps one loss value per digit per + # visit, so each sample has a loss history you can inspect in the Studio. loss_batch_mlt = criterion_mlt( preds_raw.float(), labels.long(), @@ -156,9 +272,16 @@ def train(loader, model, optimizer, criterion_mlt, device): return total_loss.detach().cpu().item() -def test(loader, model, criterion_mlt, metric_mlt, device, test_loader_len): - """Full evaluation pass over the test loader.""" +def test(loader, model, criterion_mlt, metric_mlt, device, test_loader_len, split="test"): + """Full evaluation pass over one evaluation loader. + + The metric is reset first: torchmetrics accumulates across `update` calls, + so without this each evaluation would report a running average over every + previous evaluation instead of the current one -- and the val and test + loaders would pollute each other's number. + """ losses = torch.tensor(0.0, device=device) + metric_mlt.reset() for (inputs, ids, labels) in loader: with wl.guard_testing_context: @@ -191,8 +314,8 @@ def test(loader, model, criterion_mlt, metric_mlt, device, test_loader_len): # Log per-sample metric alongside signals; persists via the storer signals = { - "test_metric/Accuracy_per_sample": acc_per_sample, - "test_metric/Inverse_Accuracy_per_sample": acc_reversed_per_sample, + f"{split}_metric/Accuracy_per_sample": acc_per_sample, + f"{split}_metric/Inverse_Accuracy_per_sample": acc_reversed_per_sample, } wl.save_signals( preds_raw=outputs, @@ -208,6 +331,298 @@ def test(loader, model, criterion_mlt, metric_mlt, device, test_loader_len): return loss.detach().cpu().item(), metric.detach().cpu().item() +def _set_dropout(model, enabled: bool) -> None: + """Toggle dropout between phases without rebuilding the model. + + Rebuilding would create a NEW watched model (new hash, new checkpoint + lineage) and lose the step-0 weights we are about to restore, so the rates + are set in place on the existing modules instead. + """ + for name, p in (("drop1", 0.25), ("drop2", 0.5)): + mod = getattr(model, name, None) + if mod is not None: + mod.p = p if enabled else 0.0 + + +def _reload_initial_weights(wl, model, optimizer, init_state, log_dir) -> dict: + """Put the model back to its step-0 random initialisation. + + Prefers WeightsLab's own step-0 checkpoint (what "reload at model age 0" + means from the Studio); falls back to the in-memory snapshot taken before + phase A. Either way the result is verified against that snapshot, so the + number reported is a measurement and not a claim. + """ + source = "in-memory snapshot" + try: + from weightslab.backend.ledgers import get_checkpoint_manager + cm = get_checkpoint_manager() + exp_hash = getattr(cm, "current_exp_hash", None) or cm.get_latest_hash() + cm.load_checkpoint(exp_hash=exp_hash, load_model=False, load_weights=True, + load_config=False, load_data=False, target_step=0, force=True) + source = "step-0 checkpoint" + except Exception as exc: + print(f"[reset] step-0 checkpoint reload unavailable ({exc}); using the snapshot", flush=True) + + def _drift(): + cur = model.state_dict() + return max((cur[k].float() - v.float()).abs().max().item() + for k, v in init_state.items() if k in cur and torch.is_tensor(cur[k])) + + worst = _drift() + if worst > 0: + model.load_state_dict(init_state, strict=False) # snapshot is authoritative + after = _drift() + source += f" + snapshot (checkpoint differed by {worst:.3g})" + worst = after + + # Adam moments must go too: keeping them would carry phase A's gradient + # history into a run that is supposed to start from scratch. + optimizer.state = type(optimizer.state)() + return {"source": source, "max_abs_diff": worst} + + +def _fmt_si(n: float, unit: str) -> str: + """1199882 -> '1.2M'. Model cards quote these rounded, so match that.""" + for div, suf in ((1e9, "G"), (1e6, "M"), (1e3, "K")): + if n >= div: + return f"{n / div:.1f}{suf}{unit}" + return f"{n:.0f}{unit}" + + +def _card(title, acc, params, n_data) -> None: + """The four-field summary, one per experiment. + + Parameters and FLOPs are identical for both runs by construction -- same + architecture, same input size. That is the point: only NumberTrainingData + moves, so any accuracy difference is attributable to the data, not capacity. + """ + print(f" {title}") + print(f" Accuracy: {acc:.2f}%") + print(f" Parameters: {_fmt_si(params, '')}") + print(f" NumberTrainingData: {n_data:,}") + + + +def _summarise(results, n_train, goldset_size) -> None: + """Print the comparison the experiment exists to produce.""" + a, b = results["signal"], results["goldset"] + macs = results["flops_fwd_per_sample"] + params = results["model_params"] + fpv = macs * 3 # fwd + bwd, per sample-visit, in MACs + + print("\n" + "=" * 74) + print(" MODEL CARDS") + print("=" * 74) + _card("full trainset (phase A: signal run)", + a["best"].get("test_acc", float("nan")), params, n_train) + print() + _card("goldset (phase B: retrained from step-0 weights)", + b["best"].get("test_acc", float("nan")), params, goldset_size) + + print("\n" + "=" * 74) + print(" TRAINING COST (what actually differs)") + print("=" * 74) + print(f"{'':24s}{'full trainset':>18s}{'goldset':>18s}") + rows = [ + ("training digits", f"{n_train:,}", f"{goldset_size:,}"), + ("steps x batch", f"{a['steps']:,} x {a['batch_size']}", f"{b['steps']:,} x {b['batch_size']}"), + ("sample visits", f"{a['sample_visits']:,}", f"{b['sample_visits']:,}"), + ("epochs over own set", f"{a['sample_visits'] / max(1, n_train):.1f}", + f"{b['sample_visits'] / max(1, goldset_size):.1f}"), + ("train TFLOPs", f"{a['sample_visits'] * fpv * 2 / 1e12:.2f}", + f"{b['sample_visits'] * fpv * 2 / 1e12:.2f}"), + ("TFLOPs to best val", f"{a['best'].get('tflops', 0) * 2:.2f}", + f"{b['best'].get('tflops', 0) * 2:.2f}"), + ("train seconds", f"{a['train_seconds']:.1f}", f"{b['train_seconds']:.1f}"), + ("eval seconds", f"{a['eval_seconds']:.1f}", f"{b['eval_seconds']:.1f}"), + ("wall seconds", f"{a['wall_seconds']:.1f}", f"{b['wall_seconds']:.1f}"), + ("best val acc %", f"{a['best'].get('val_acc', float('nan')):.2f}", + f"{b['best'].get('val_acc', float('nan')):.2f}"), + ("test @ best val %", f"{a['best'].get('test_acc', float('nan')):.2f}", + f"{b['best'].get('test_acc', float('nan')):.2f}"), + ("step of best val", f"{a['best'].get('step', 0):,}", f"{b['best'].get('step', 0):,}"), + ] + for label, av, bv in rows: + print(f"{label:24s}{av:>18s}{bv:>18s}") + + ep_a = n_train * fpv * 2 / 1e12 + ep_b = goldset_size * fpv * 2 / 1e12 + saved = f"{(1 - ep_b / ep_a) * 100:.1f}% less per epoch" if ep_a > 0 else "n/a" + print("\n one epoch over its own training set:") + print(f" full trainset: {ep_a:.3f} TFLOPs over {n_train:,} digits") + print(f" goldset: {ep_b:.3f} TFLOPs over {goldset_size:,} digits ({saved})") + print(" Same architecture, same per-sample FLOPs -- training cost scales purely") + print(" with sample-visits, so the saving is in the data, not the model.") + print("=" * 74) + + +# ============================================================================ +# Experiment phases: signal run -> goldset -> retrain from the initial weights +# ============================================================================ +def model_flops_per_sample(model, device) -> int: + """Forward MACs for one sample, counted from the conv/linear layers. + + The model is identical in both phases, so this is a constant; what differs + between the phases is only how many sample-visits each one pays for. + """ + macs = {"n": 0} + hooks = [] + + def conv_hook(mod, inp, out): + macs["n"] += out.numel() * mod.in_channels // mod.groups * mod.kernel_size[0] * mod.kernel_size[1] + + def lin_hook(mod, inp, out): + macs["n"] += out.numel() * mod.in_features + + for m in model.modules(): + if isinstance(m, nn.Conv2d): + hooks.append(m.register_forward_hook(conv_hook)) + elif isinstance(m, nn.Linear): + hooks.append(m.register_forward_hook(lin_hook)) + was_training = model.training + model.eval() + with torch.no_grad(): + model(torch.zeros(1, 1, 28, 28, device=device)) + model.train(was_training) + for h in hooks: + h.remove() + if macs["n"] == 0: + raise RuntimeError( + "FLOPs count came back 0 -- no Conv2d/Linear was visited. Pass the " + "UNWRAPPED model (the object built before wl.watch_or_edit): a watched " + "model is a ModelInterface whose .modules() yields only itself.") + return int(macs["n"]) + + +def build_goldset(wl, steps_per_pass, labels_by_id, pool_ids, gcfg): + """Per class, the digits with the highest mean loss over passes 2..N, minus + the noisy ones (median loss over all passes above the threshold). + + Reads the per-sample loss history WeightsLab recorded during the signal run + -- no second inference pass, no extra bookkeeping of our own. + """ + wl.drain_signals() + rows = wl.query_signal_history("train-loss-CE") + hist = pd.DataFrame(rows, columns=["sample_id", "step", "loss", "run"][:len(rows[0])]) if rows else pd.DataFrame() + if hist.empty: + raise SystemExit("[goldset] no train-loss-CE history recorded; cannot build the goldset") + hist["sample_id"] = hist["sample_id"].astype(int) + hist["loss"] = hist["loss"].astype(float) + # "Pass n" is the n-th time THIS digit was seen, taken from its own visit + # order rather than from step arithmetic: the recorded step numbering is not + # guaranteed to start at 1, and a uniform shuffle is what makes the two + # equivalent in the first place. + hist = hist.sort_values(["sample_id", "step"], kind="stable") + hist["pass"] = hist.groupby("sample_id").cumcount() + per_pass = hist.groupby(["sample_id", "pass"])["loss"].last().unstack() + + noisy = per_pass.median(axis=1, skipna=True) > float(gcfg.get("noisy_median_loss", 1.0)) + first = 1 if gcfg.get("skip_first_pass", True) else 0 + n_passes = int(per_pass.shape[1]) + if n_passes <= first: + raise SystemExit( + f"[goldset] only {n_passes} pass(es) recorded; the rule scores passes " + f"{first + 1}..N, so it needs at least {first + 2}. Per-sample loss over " + f"fewer passes measures batch placement, not difficulty.") + score = per_pass.iloc[:, first:].mean(axis=1, skipna=True) + + pool = set(int(i) for i in pool_ids) + cand = pd.DataFrame({"score": score, "noisy": noisy.reindex(score.index, fill_value=False)}) + cand = cand[cand.index.map(lambda i: int(i) in pool) & ~cand["noisy"] & cand["score"].notna()] + cand["label"] = [labels_by_id[int(i)] for i in cand.index] + + k = int(gcfg.get("per_class", 360)) + gold, per_class_got = [], {} + for c in sorted(cand["label"].unique()): + g = cand[cand["label"] == c].sort_values("score", ascending=False, kind="stable") + picked = [int(i) for i in g.head(k).index] + gold += picked + per_class_got[int(c)] = len(picked) + + # A goldset that is empty, or short of its per-class quota, means the rule or + # the history is wrong -- stop rather than retrain on a malformed subset. + if not gold: + raise SystemExit("[goldset] selection is empty; refusing to continue") + short = {c: n for c, n in per_class_got.items() if n < k} + if short: + raise SystemExit(f"[goldset] classes below the {k}/class quota: {short}") + if len(set(gold)) != len(gold): + raise SystemExit("[goldset] duplicate ids in the selection") + + return sorted(gold), {"noisy": int(noisy.sum()), "scored": int(len(cand)), + "passes": int(per_pass.shape[1]), "per_class": per_class_got} + + +def run_phase(name, total_steps, recipe, eval_every, ctx): + """Train one phase, evaluating validation on a cadence and TEST ONLY when + validation reaches a new maximum. + + Test accuracy is therefore never used to steer anything -- it is read at + exactly the points where validation says the model just got better, which + is the only honest moment to look at it. + """ + wl, model, optimizer = ctx["wl"], ctx["model"], ctx["optimizer"] + train_loader, val_loader, test_loader = ctx["train_loader"], ctx["val_loader"], ctx["test_loader"] + device = ctx["device"] + + flops_per_visit = ctx["flops_fwd_per_sample"] * 3 # fwd + bwd ~= 3x fwd + bs = recipe["batch_size"] + + best_val, best = -1.0, {} + train_seconds = eval_seconds = 0.0 + visits = 0 + history = [] + t_phase = time.perf_counter() + + for step in range(1, total_steps + 1): + t0 = time.perf_counter() + loss = train(train_loader, model, optimizer, ctx["train_criterion"], device, step, recipe) + train_seconds += time.perf_counter() - t0 + visits += bs + + if step % eval_every == 0 or step == total_steps: + t0 = time.perf_counter() + val_loss, val_acc = test(val_loader, model, ctx["val_criterion"], ctx["val_metric"], + device, ctx["val_loader_len"], split="val") + row = {"phase": name, "step": step, "train_loss": round(loss, 5), + "val_acc": round(val_acc, 4), "test_acc": None, + "visits": visits, "tflops": round(visits * flops_per_visit / 1e12, 4), + "train_s": round(train_seconds, 1)} + signals = {"val/accuracy": val_acc} + + if val_acc > best_val: # new best validation -> read test once + best_val = val_acc + _, test_acc = test(test_loader, model, ctx["test_criterion"], ctx["test_metric"], + device, ctx["test_loader_len"], split="test") + row["test_acc"] = round(test_acc, 4) + signals["test/accuracy"] = test_acc + best = {"step": step, "val_acc": val_acc, "test_acc": test_acc, + "visits": visits, "tflops": row["tflops"], + "train_s": round(train_seconds, 1)} + eval_seconds += time.perf_counter() - t0 + + wl.save_model_signals(signals) + history.append(row) + mark = " <- new best val, test read" if row["test_acc"] is not None else "" + print(f"[{name}] step {step:5d}/{total_steps} loss {loss:.4f} " + f"val {val_acc:6.2f}% test {row['test_acc'] if row['test_acc'] is not None else ' - '}" + f" {row['tflops']:.2f} TFLOPs{mark}", flush=True) + + return { + "phase": name, + "steps": total_steps, + "batch_size": bs, + "sample_visits": visits, + "train_seconds": round(train_seconds, 1), + "eval_seconds": round(eval_seconds, 1), + "wall_seconds": round(time.perf_counter() - t_phase, 1), + "flops_fwd_per_sample": ctx["flops_fwd_per_sample"], + "train_tflops": round(visits * flops_per_visit / 1e12, 4), + "best": best, + "history": history, + } + + # ----------------------------------------------------------------------------- # Main # ----------------------------------------------------------------------------- @@ -227,6 +642,20 @@ def test(loader, model, criterion_mlt, metric_mlt, device, test_loader_len): parameters.setdefault("device", "auto") parameters.setdefault("training_steps_to_do", 1000000) parameters.setdefault("eval_full_to_steps_ratio", 50) + parameters.setdefault("seed", 0) + parameters.setdefault("augment", False) + parameters.setdefault("dropout", False) + parameters.setdefault("optimizer", {}).setdefault("schedule", "const") + parameters["optimizer"].setdefault("weight_decay", 0.0) + + # Deterministic seeding: a rerun reproduces the same shuffle, and therefore + # the same per-digit visit order. + seed = int(parameters["seed"]) + random.seed(seed) + np.random.seed(seed) + torch.manual_seed(seed) + torch.backends.cudnn.deterministic = True + torch.backends.cudnn.benchmark = False # Experiment name exp_name = parameters["experiment_name"] @@ -259,12 +688,14 @@ def test(loader, model, criterion_mlt, metric_mlt, device, test_loader_len): enable_h5_persistence = parameters.get("enable_h5_persistence", True) # Model - _model = CNN().to(device) + _model = CNN(dropout=bool(parameters["dropout"])).to(device) model = wl.watch_or_edit(_model, flag="model", device=device) # Optimizer - lr = parameters.get("optimizer", {}).get("lr", 0.01) - _optimizer = optim.Adam(model.parameters(), lr=lr) + opt_cfg = parameters["optimizer"] + lr = opt_cfg.get("lr", 0.001) + _optimizer = optim.Adam( + model.parameters(), lr=lr, weight_decay=float(opt_cfg["weight_decay"])) optimizer = wl.watch_or_edit( _optimizer, flag="optimizer", @@ -286,9 +717,10 @@ def test(loader, model, criterion_mlt, metric_mlt, device, test_loader_len): # Read data config for all loaders train_cfg = parameters.get("data", {}).get("train_loader", {}) + val_cfg = parameters.get("data", {}).get("val_loader", {}) test_cfg = parameters.get("data", {}).get("test_loader", {}) - _train_dataset = MNISTCustomDataset( + _full_train = MNISTCustomDataset( root=data_root, train=True, download=should_download, @@ -311,7 +743,35 @@ def test(loader, model, criterion_mlt, metric_mlt, device, test_loader_len): max_samples=test_cfg.get("max_samples", None) ) - # Create tracked loaders for train, test, and test + # Split the 60,000 MNIST train images into the training pool and a held-out + # validation set BEFORE anything is registered: the two loaders then own + # disjoint images, so nothing has to be discarded afterwards and no image is + # registered twice. + # + # WeightsLab numbers each split's samples contiguously from where the + # previous one ended, so a sample id here is a position in the split, not an + # MNIST index. Both index lists are written next to this config, in loader + # order, so the mapping back is a lookup: + # + # mnist_index = train_ids[wl_sample_id] (train split) + # mnist_index = val_ids[wl_sample_id - len(train_ids)] (val split) + val_ids_path = parameters.get("val_ids") + if val_ids_path and not os.path.isabs(val_ids_path): + val_ids_path = os.path.join(os.path.dirname(__file__), val_ids_path) + val_ids = get_val_ids(_full_train.mnist.targets, val_ids_path) + held_out = set(val_ids) + train_ids = [i for i in range(len(_full_train)) if i not in held_out] + assert not (set(train_ids) & held_out), "train ids overlap the validation split" + + if val_ids_path: # same order the loader serves them in + with open(os.path.join(os.path.dirname(val_ids_path), "train_ids.txt"), "w") as fh: + fh.write("\n".join(map(str, train_ids)) + "\n") + + _train_dataset = SubsetWithIds(_full_train, train_ids) + _val_dataset = SubsetWithIds(_full_train, val_ids) + n_train = len(_train_dataset) + + # Create tracked loaders for train, val and test train_loader = wl.watch_or_edit( _train_dataset, flag="data", @@ -324,6 +784,18 @@ def test(loader, model, criterion_mlt, metric_mlt, device, test_loader_len): preload_metadata=False, enable_h5_persistence=enable_h5_persistence ) + val_loader = wl.watch_or_edit( + _val_dataset, + flag="data", + loader_name="val_loader", + batch_size=val_cfg.get("batch_size", 500), + shuffle=val_cfg.get("shuffle", False), + is_training=False, + compute_hash=False, + preload_labels=True, + preload_metadata=False, + enable_h5_persistence=enable_h5_persistence + ) test_loader = wl.watch_or_edit( _test_dataset, flag="data", @@ -337,104 +809,142 @@ def test(loader, model, criterion_mlt, metric_mlt, device, test_loader_len): enable_h5_persistence=enable_h5_persistence ) + # 8 passes over 55,000 digits at batch 64 = 8 x 860 = 6,880 steps. + batch_size = train_cfg.get("batch_size", 16) + steps_per_pass = -(-n_train // batch_size) + if parameters.get("epochs"): + parameters["training_steps_to_do"] = int(parameters["epochs"]) * steps_per_pass + total_steps = int(parameters["training_steps_to_do"]) + recipe = { + "augment": bool(parameters["augment"]), + "lr": float(lr), + "schedule": opt_cfg["schedule"], + "total_steps": total_steps, + } + # Losses & metrics (watched objects – they log themselves) train_criterion = wl.watch_or_edit( nn.CrossEntropyLoss(reduction="none"), - flag="loss", signal_name="train-loss-CE", log=True) + flag="loss", signal_name="train-loss-CE", per_sample=True, log=True) test_criterion = wl.watch_or_edit( nn.CrossEntropyLoss(reduction="none"), flag="loss", signal_name="test-loss-CE", log=True) + val_criterion = wl.watch_or_edit( + nn.CrossEntropyLoss(reduction="none"), + flag="loss", signal_name="val-loss-CE", log=True) + metric = wl.watch_or_edit( Accuracy(task="multiclass", num_classes=10).to(device), flag="metric", signal_name="metric-ACC", log=True) + val_metric = wl.watch_or_edit( + Accuracy(task="multiclass", num_classes=10).to(device), + flag="metric", signal_name="metric-ACC-val", log=True) # Start WeightsLab services (gRPC only, no CLI) wl.serve( serving_grpc=parameters.get("serving_grpc", False) ) - print("=" * 60) - print(" STARTING TRAINING") - print(f" Evaluation every {eval_full_to_train_steps_ratio} steps") - print(f" Dataset splits: train={len(_train_dataset)}, test={len(_test_dataset)}") - print(f" Logs will be saved to: {log_dir}") - print("=" * 60 + "\n") - - # Setup clean progress bar with custom format - # Training runs until YOU stop it -- from the studio's pause button, the CLI, - # or Ctrl+C. itertools.count() rather than range(training_steps_to_do): a - # predefined step budget ends the process mid-experiment, which is the - # opposite of how WeightsLab is used (inspect the curves, edit the data or - # the architecture, keep going). `training_steps_to_do` remains a live - # hyperparameter for the UI's own "run N more steps" control; it is not a - # ceiling on this loop. - if tqdm_display: - train_range = tqdm.tqdm( - itertools.count(), - desc="Training", - bar_format="{desc}: {n} steps [{elapsed}, {rate_fmt}] {bar} | {postfix}", - ncols=140, - position=0, - leave=True - ) - else: - train_range = itertools.count() - - # ============= - # Training Loop - wl.start_training(timeout=3) # Blocks and keeps the main thread alive while background services run. Optionally set a timeout (seconds) to auto-stop. - - train_loss = None - test_loss, test_metric = None, None - test_loader_len = len(test_loader) # Store length before wrapping with tqdm - for train_step in train_range: - age = model.get_age() if hasattr(model, "get_age") else train_step # Get model age in steps (not necessarily equal to train_step if model was reloaded or has seen more data than training steps) - - # Train one step - train_loss = train(train_loader, model, optimizer, train_criterion, device) - - # Periodic test evaluation - if age > 0 and age % eval_full_to_train_steps_ratio == 0: - # Test (no nested progress bar) - test_loss, test_metric = test( - test_loader, - model, - test_criterion, - metric, - device, - test_loader_len - ) + # ---- constants shared by both phases ------------------------------------- + val_loader_len = len(val_loader) + test_loader_len = len(test_loader) + flops_fwd = model_flops_per_sample(_model, device) # unwrapped: see the docstring + labels_by_id = {i: int(_full_train.mnist.targets[m]) for i, m in enumerate(train_ids)} + gcfg = parameters.get("goldset", {}) or {} + + ctx = dict(wl=wl, model=model, optimizer=optimizer, device=device, + train_loader=train_loader, val_loader=val_loader, test_loader=test_loader, + train_criterion=train_criterion, val_criterion=val_criterion, + test_criterion=test_criterion, val_metric=val_metric, test_metric=metric, + val_loader_len=val_loader_len, test_loader_len=test_loader_len, + flops_fwd_per_sample=flops_fwd) + + p_sig = parameters["phases"]["signal"] + p_gold = parameters["phases"]["goldset"] + sig_steps = int(p_sig["epochs"]) * (-(-n_train // int(p_sig["batch_size"]))) + + print("=" * 72) + print(" TWO-PHASE EXPERIMENT (one process, one experiment directory)") + print(f" splits: train={n_train} val={len(_val_dataset)} test={len(_test_dataset)}") + print(f" model: {sum(p.numel() for p in model.parameters()):,} params, " + f"{flops_fwd / 1e6:.2f} MFLOPs forward per sample ({flops_fwd * 3 / 1e6:.2f} incl. backward)") + print(f" phase A (signal): {sig_steps} steps x batch {p_sig['batch_size']} " + f"= {p_sig['epochs']} passes over {n_train} digits") + print(f" phase B (goldset): {p_gold['training_steps_to_do']} steps x batch {p_gold['batch_size']}") + print(f" logs: {log_dir}") + print("=" * 72 + "\n") + + # Hand control to WeightsLab before the first guarded step: without this the + # training guard never opens and the loop stalls with the process alive. + wl.start_training(timeout=int(parameters.get("start_training_timeout", 1))) + + results = {"splits": {"train": n_train, "val": len(_val_dataset), "test": len(_test_dataset)}, + "model_params": int(sum(p.numel() for p in model.parameters())), + "flops_fwd_per_sample": flops_fwd} + t_all = time.perf_counter() + + # ---- the initial random weights, kept so phase B can start from them ------ + # Tensors only: a watched model's state_dict also carries scalar bookkeeping + # (the age counter), which has nothing to clone or compare. + init_state = {k: v.detach().clone() for k, v in model.state_dict().items() + if torch.is_tensor(v)} + + # ---- PHASE A: signal run over the whole training pool --------------------- + recipe_sig = {"augment": bool(p_sig["augment"]), "lr": float(p_sig["lr"]), + "schedule": p_sig["schedule"], "total_steps": sig_steps, + "batch_size": int(p_sig["batch_size"])} + train_loader.set_batch_size(int(p_sig["batch_size"])) + _set_dropout(model, bool(p_sig["dropout"])) + for g in optimizer.param_groups: + g["lr"] = float(p_sig["lr"]) + results["signal"] = run_phase("signal", sig_steps, recipe_sig, int(p_sig["eval_every"]), ctx) + + # ---- build the goldset from what phase A recorded ------------------------- + t0 = time.perf_counter() + steps_per_pass = -(-n_train // int(p_sig["batch_size"])) + goldset, ginfo = build_goldset(wl, steps_per_pass, labels_by_id, range(n_train), gcfg) + wl.tag_samples(goldset, "goldset") + results["goldset_build"] = {**ginfo, "size": len(goldset), + "seconds": round(time.perf_counter() - t0, 1)} + mnist_ids = sorted(int(train_ids[i]) for i in goldset) + with open(os.path.join(log_dir, "goldset_ids_mnist.txt"), "w") as fh: + fh.write("\n".join(map(str, mnist_ids)) + "\n") + print(f"\n[goldset] {len(goldset)} digits ({len(goldset) / n_train:.2%} of the pool) " + f"from {ginfo['passes']} passes; {ginfo['noisy']} noisy excluded; " + f"tagged 'goldset'; MNIST ids -> goldset_ids_mnist.txt\n", flush=True) + + # ---- reload the INITIAL weights (model age 0), then train on the goldset --- + reload_info = _reload_initial_weights(wl, model, optimizer, init_state, log_dir) + results["reload"] = reload_info + print(f"[reset] restored the step-0 weights ({reload_info['source']}); " + f"max|w - w0| = {reload_info['max_abs_diff']}\n", flush=True) + + keep = set(int(i) for i in goldset) + wl.discard_samples([i for i in range(n_train) if i not in keep]) + + recipe_gold = {"augment": bool(p_gold["augment"]), "lr": float(p_gold["lr"]), + "schedule": p_gold["schedule"], "total_steps": int(p_gold["training_steps_to_do"]), + "batch_size": int(p_gold["batch_size"])} + train_loader.set_batch_size(int(p_gold["batch_size"])) + _set_dropout(model, bool(p_gold["dropout"])) + results["goldset"] = run_phase("goldset", int(p_gold["training_steps_to_do"]), + recipe_gold, int(p_gold["eval_every"]), ctx) + + # ---- report --------------------------------------------------------------- + results["wall_seconds_total"] = round(time.perf_counter() - t_all, 1) + _summarise(results, n_train, len(goldset)) + with open(os.path.join(log_dir, "experiment_results.json"), "w") as fh: + json.dump(results, fh, indent=1) + print(f"\n results -> {os.path.join(log_dir, 'experiment_results.json')}") - # Verbose - if verbose and not tqdm_display: - import sys - # Build compact progress message - msg = f"Step {train_step} (Age {age}): Loss={train_loss:.4f}" - if test_loss is not None: - msg += f" | Test={test_loss:.4f} ({test_metric:.1f}%)" - - # Clear line completely and print (pad to 100 chars to overwrite previous content) - sys.stdout.write(f"\r{msg:<100}") - sys.stdout.flush() - elif tqdm_display: - # Build compact postfix string - postfix_parts = [f"train_loss={train_loss:.4f}"] - if test_loss is not None: - postfix_parts.append(f"test_loss={test_loss:.4f}") - if test_metric is not None: - postfix_parts.append(f"test_acc={test_metric:.1f}%") - - train_range.set_postfix_str(" | ".join(postfix_parts)) - - print("\n" + "=" * 60) - print(f" Training completed in {time.time() - start_time:.2f} seconds") - print(f" Logs saved to: {log_dir}") - print("=" * 60) - - # Final export of signal history and data grid to root_log_dir wl.write_history() wl.write_dataframe() - # Keep the main thread alive to allow background serving threads to run - wl.keep_serving() + # Keep the main thread alive so the Studio stays attached. Set + # keep_serving: false to exit once the results are written -- a batch sweep + # runs these back to back and must not block on the last one. + if parameters.get("keep_serving", True): + wl.keep_serving() + else: + print(" keep_serving: false -> exiting", flush=True) From b2c23fa3ad0481136da3e6146f740b517ea2ecc7 Mon Sep 17 00:00:00 2001 From: GuillaumePELLUET Date: Sat, 26 Sep 2026 00:58:40 +0200 Subject: [PATCH 16/29] fix chkpt restoring from several root dirs --- weightslab/components/checkpoint_manager.py | 88 ++++++++++++++++--- weightslab/data/dataframe_manager.py | 82 ++++++++++++++++- .../trainer/services/experiment_service.py | 9 ++ weightslab/utils/logs.py | 30 +++---- 4 files changed, 179 insertions(+), 30 deletions(-) diff --git a/weightslab/components/checkpoint_manager.py b/weightslab/components/checkpoint_manager.py index fbbab5e2..ca03b640 100644 --- a/weightslab/components/checkpoint_manager.py +++ b/weightslab/components/checkpoint_manager.py @@ -606,13 +606,15 @@ def _extract_step_from_checkpoint_name(self, filename: str) -> Optional[int]: except Exception: return None - def _select_weight_checkpoint_file(self, exp_hash: str, target_step: Optional[int] = None) -> Optional[Path]: + def _select_weight_checkpoint_file(self, exp_hash: str, target_step: Optional[int] = None, + models_dir: Optional[Path] = None) -> Optional[Path]: """Select weight checkpoint file for an experiment hash. - If target_step is None: returns latest checkpoint. - If target_step is provided: returns closest step; tie breaks toward higher step. + - models_dir: where to look (default: this root's; a sibling root's in multi-root mode). """ - model_dir = self.models_dir / exp_hash[8:-8] + model_dir = (models_dir or self.models_dir) / exp_hash[8:-8] if not model_dir.exists(): return None @@ -1718,6 +1720,35 @@ def _load_manifest(self) -> Dict[str, Any]: logger.warning(f"Failed to load manifest: {e}") return {'experiments': {}, 'latest_hash': None} + def _locate_experiment(self, exp_hash: str) -> Optional[tuple]: + """Find the checkpoint root that owns *exp_hash*. + + Returns ``(checkpoints_dir, manifest)``: the effective root's when its + manifest lists the hash (the usual, single-root case). In multi-root mode + the other discovered roots only contribute curves to the logger, while + their checkpoints stay in their own directories; when the hash belongs to + one of them, return that root's checkpoints directory and manifest so the + hash can still be restored. If several siblings list the same hash (e.g. + copies of one run), the most recently updated one wins, consistent with + :meth:`_resolve_effective_root_dir`. ``None`` if no discovered root has it. + """ + manifest = self._load_manifest() + if exp_hash in (manifest.get('experiments') or {}): + return self.checkpoints_dir, manifest + + siblings = [d for d in self.discovered_root_dirs if d != self.root_log_dir] + ranked = sorted(siblings, key=lambda d: (self._read_root_manifest_timestamp(d), str(d)), reverse=True) + for root in ranked: + try: + with open(root / "checkpoints" / "manifest.yaml", 'r') as f: + sibling_manifest = yaml.safe_load(f) or {} + except Exception as e: + logger.debug(f"Skipping unreadable manifest under sibling root {root}: {e}") + continue + if exp_hash in (sibling_manifest.get('experiments') or {}): + return root / "checkpoints", sibling_manifest + return None + def _load_manager_state(self): """Load manager state if available""" state_file = self.root_log_dir / ".checkpoint_manager_state.json" @@ -1948,6 +1979,8 @@ def load_checkpoint(self, - 'weights': Checkpoint dict with weights and metadata - 'config': Loaded config (if changed and load_config=True) - 'data_state': Loaded data state (if changed and load_data=True) + - 'data_store': Data store (data.h5) of the sibling root owning the + hash, in multi-root mode (if load_data=True) - 'loaded_components': Set of components that were loaded - 'exp_hash': The experiment hash that was loaded """ @@ -1956,16 +1989,24 @@ def load_checkpoint(self, 'weights': None, 'config': None, 'data_state': None, + 'data_store': None, 'rng_state': None, 'loaded_components': set(), 'exp_hash': exp_hash } - # Load manifest to get component hashes - manifest = self._load_manifest() - if exp_hash not in manifest.get('experiments', {}): + # Find the root that owns this hash: this root, or in multi-root mode the + # sibling root it came from. Everything below is read from that root. + located = self._locate_experiment(exp_hash) + if located is None: logger.error(f"Experiment hash {exp_hash} not found in manifest") return result + checkpoints_dir, manifest = located + models_base = checkpoints_dir / "models" + hp_base = checkpoints_dir / "HP" + data_base = checkpoints_dir / "data" + if checkpoints_dir != self.checkpoints_dir: + logger.info(f"Hash {exp_hash[:16]} belongs to sibling root {checkpoints_dir.parent}; loading it from there") exp_info = manifest['experiments'][exp_hash] target_hp_hash = exp_info.get('hp_hash') target_model_hash = exp_info.get('model_hash') @@ -1985,7 +2026,7 @@ def load_checkpoint(self, # Load model architecture if different, or load only RNG state for reproducibility if model hash is unchanged model_rng_loaded = False if load_model and (target_model_hash != current_model_hash or force): - model_dir = self.models_dir / exp_hash[8:-8] + model_dir = models_base / exp_hash[8:-8] arch_ref_file = model_dir / f"{exp_hash[8:-8]}_architecture_ref.json" # First check if this is a reference to architecture @@ -2000,7 +2041,7 @@ def load_checkpoint(self, logger.warning(f"Failed to load architecture reference: {e}") # Now load from actual location - actual_arch_file = self.models_dir / actual_arch_hash / f"{actual_arch_hash}_architecture.pkl" + actual_arch_file = models_base / actual_arch_hash / f"{actual_arch_hash}_architecture.pkl" if actual_arch_file.exists(): try: @@ -2022,7 +2063,7 @@ def load_checkpoint(self, elif load_model and (target_model_hash == current_model_hash and not force): # Try to load only the RNG state from the latest model checkpoint for reproducibility - model_dir = self.models_dir / exp_hash[8:-8] + model_dir = models_base / exp_hash[8:-8] checkpoint_files = sorted(model_dir.glob(f"{exp_hash}_step_*.pt")) if not checkpoint_files: checkpoint_files = sorted(model_dir.glob(f"{exp_hash[8:-8]}_step_*.pt")) @@ -2045,7 +2086,7 @@ def load_checkpoint(self, # Load model weights (always if requested) if load_weights: - model_dir = self.models_dir / exp_hash[8:-8] + model_dir = models_base / exp_hash[8:-8] # First, try to get the weight checkpoint from manifest for this specific experiment checkpoint_file_to_load = None @@ -2060,7 +2101,7 @@ def load_checkpoint(self, # Fallback: scan for weight files (old behavior for backward compatibility) if checkpoint_file_to_load is None: - checkpoint_file_to_load = self._select_weight_checkpoint_file(exp_hash, target_step=target_step) + checkpoint_file_to_load = self._select_weight_checkpoint_file(exp_hash, target_step=target_step, models_dir=models_base) if checkpoint_file_to_load is not None: if target_step is None: logger.debug(f" Using latest weight checkpoint from directory scan: {checkpoint_file_to_load.name}") @@ -2105,7 +2146,7 @@ def load_checkpoint(self, # Load config if different if load_config and (target_hp_hash != current_hp_hash or force): - hp_dir = self.hp_dir / exp_hash[:8] + hp_dir = hp_base / exp_hash[:8] config_file = hp_dir / f"{exp_hash[:8]}_config.yaml" if config_file.exists(): @@ -2124,7 +2165,7 @@ def load_checkpoint(self, # Load data snapshot if different, or if only RNG state changed (for reproducibility) if load_data: - data_dir = self.data_checkpoint_dir / exp_hash[-8:] + data_dir = data_base / exp_hash[-8:] json_file = data_dir / f"{exp_hash[-8:]}_data_snapshot.json" # Always try to load RNG state for reproducibility, even if data hash is unchanged @@ -2162,6 +2203,12 @@ def load_checkpoint(self, else: logger.warning(f" [WARNING] Data snapshot file not found: {json_file}") + # The snapshot only holds sample ids, tags and discards. The rest of a + # sibling root's per-sample stats (last signal values, predictions...) + # lives in its own data store, which this root never loaded. + if checkpoints_dir != self.checkpoints_dir and (data_base / "data.h5").exists(): + result['data_store'] = data_base / "data.h5" + logger.info(f"Loaded components: {result['loaded_components']}") return result @@ -2331,6 +2378,17 @@ def load_state( logger.error(f"[ERROR] Failed to apply config: {e}") self.error_loading_checkpoint.append('config') if 'config' not in self.error_loading_checkpoint else None # Reset first_time to allow future auto-resume attempts if config application failed + # Multi-root: load the sibling run's per-sample stats first, so the + # checkpoint's own tags and discards (the snapshot below) apply on top + if checkpoint_data.get('data_store') is not None: + try: + dfm = ledgers.get_dataframe() + if dfm != None: + rows = dfm.import_from_store(checkpoint_data['data_store']) + logger.info(f"[OK] Loaded per-sample stats from {checkpoint_data['data_store']} ({rows} rows)") + except Exception as e: + logger.warning(f"[WARNING] Could not load per-sample stats from {checkpoint_data['data_store']}: {e}") + # Apply data (merge snapshot columns into current dataframe) if 'data' in checkpoint_data['loaded_components']: try: @@ -2438,8 +2496,10 @@ def load_state( self.current_exp_hash = exp_hash self.previous_exp_hash = old_hash - # Keep hash generator in sync with loaded experiment - manifest = self._load_manifest() + # Keep hash generator in sync with loaded experiment (in multi-root + # mode its manifest may live in a sibling root) + located = self._locate_experiment(exp_hash) + manifest = located[1] if located else self._load_manifest() exp_info = manifest.get('experiments', {}).get(exp_hash, {}) component_hashes = { 'hp': exp_info.get('hp_hash'), diff --git a/weightslab/data/dataframe_manager.py b/weightslab/data/dataframe_manager.py index 064a6a7c..1cf7aeca 100644 --- a/weightslab/data/dataframe_manager.py +++ b/weightslab/data/dataframe_manager.py @@ -10,7 +10,8 @@ import torch from datetime import datetime -from typing import Dict, Sequence, Any, List +from pathlib import Path +from typing import Dict, Sequence, Any, List, Iterable, Optional from weightslab.data.h5_dataframe_store import H5DataFrameStore from weightslab.data.h5_array_store import H5ArrayStore @@ -853,6 +854,85 @@ def _load_existing_data(self, origin: str = None, autoload_arrays: bool | list | else: logger.warning(f"[LedgeredDataFrameManager] Loaded data missing 'sample_id' column for origin={origin}. Skipping load.") + def import_from_store(self, store: H5DataFrameStore | str | Path, origins: Optional[Iterable[str]] = None) -> int: + """Replace this ledger's per-sample stats with the ones persisted in + another root's H5 store -- e.g. a sibling experiment restored from a + multi-root viewer: last signal values, predictions, nb_seen, tags, + discarded. ``store`` is only read. + + Only origins already registered here are imported, and only for samples + this ledger knows. Signal and tag columns this ledger holds but + ``store`` lacks belong to another run, so they are cleared on the + imported rows, and dropped once no row holds a value. Array references + are skipped: their arrays live in the other root's arrays.h5. Returns + the number of rows imported. + """ + if not isinstance(store, H5DataFrameStore): + store = H5DataFrameStore(Path(store)) + if not store.exists(): + return 0 + + with self._lock: + if self._df.empty or SampleStats.Ex.ORIGIN.value not in self._df.columns: + return 0 + registered = {str(o) for o in self._df[SampleStats.Ex.ORIGIN.value].dropna().unique()} + known_samples = set(self._df.index.get_level_values(0)) + wanted = registered if origins is None else registered & {str(o) for o in origins} + + # Bring the full category sets of categorical tags along. + registry = store.load_tag_registry() + if registry: + with self._lock: + for name, cats in registry.items(): + self._merge_categories(name, cats) + + run_column_prefixes = ("signals", "SIGNALS", f"{SampleStats.Ex.TAG.value}:", "TAG:") + stale_columns = set() + imported = 0 + for origin in sorted(wanted): + loaded = store.load(origin) # read-only, unlike load_all() + if loaded.empty or "sample_id" not in loaded.columns: + continue + if "annotation_id" in loaded.columns: + loaded = loaded.set_index(['sample_id', 'annotation_id']) + else: + loaded = self._expand_dataframe_with_annotations(loaded.set_index("sample_id")) + loaded.index = pd.MultiIndex.from_arrays( + [pd.Index([self._normalize_sample_id(v) for v in loaded.index.get_level_values(0)]), + loaded.index.get_level_values(1)], + names=['sample_id', 'annotation_id'], + ) + loaded = loaded[loaded.index.get_level_values(0).isin(known_samples)] + if loaded.empty: + continue + + for col in [c for c in self._array_columns if c in loaded.columns]: + if loaded[col].dtype == object and loaded[col].map(lambda v: isinstance(v, str) and '.h5:/' in v).any(): + loaded = loaded.drop(columns=col) + + with self._lock: + stale = [c for c in self._df.columns + if c not in loaded.columns and str(c).startswith(run_column_prefixes)] + for col in stale: + loaded[col] = False if pd.api.types.is_bool_dtype(self._df[col].dtype) else np.nan + stale_columns.update(stale) + + self.upsert_df(loaded, origin=origin, force_flush=True) + imported += len(loaded) + + # Drop the other run's columns that no row holds anymore (a cleared tag holds False) + with self._lock: + for col in stale_columns: + if col not in self._df.columns: + continue + values = self._df[col].astype(object) + held = values.notna() + if not str(col).startswith(("signals", "SIGNALS")): + held &= values != False + if not held.any(): + self._df.pop(col) + return imported + def upsert_df(self, df_local: List | pd.DataFrame, origin: str = None, force_flush: bool = False): if df_local is None or (isinstance(df_local, pd.DataFrame) and df_local.empty) or len(df_local) == 0: return diff --git a/weightslab/trainer/services/experiment_service.py b/weightslab/trainer/services/experiment_service.py index d0c57bdd..5e0c3fdc 100644 --- a/weightslab/trainer/services/experiment_service.py +++ b/weightslab/trainer/services/experiment_service.py @@ -609,6 +609,15 @@ def RestoreCheckpoint(self, request, context): # Reply if success: logger.info(f"Successfully restored checkpoint: {experiment_hash}") + # A restore can change any per-sample column (tags, discards and, for + # a sibling root's run in multi-root mode, all of its stats): rebuild + # the grid view now; the partial refresh only syncs a few columns. + data_service = getattr(self, "data_service", None) + if data_service is not None: + try: + data_service._slowUpdateInternals(force=True) + except Exception as e: + logger.debug(f"Grid view refresh after restore failed: {e}") self._log_audit( "checkpoint_restore", "success", diff --git a/weightslab/utils/logs.py b/weightslab/utils/logs.py index 717165b4..2a5e33d8 100644 --- a/weightslab/utils/logs.py +++ b/weightslab/utils/logs.py @@ -130,7 +130,7 @@ def setup_logging(level, log_to_file=True): # File handler - write to temp directory if log_to_file: # Create temp directory for logs if it doesn't exist - temp_dir = tempfile.mkdtemp() + temp_dir = tempfile.mkdtemp() if not os.environ.get('WEIGHTSLAB_ROOT_LOG_DIR') else os.environ.get('WEIGHTSLAB_ROOT_LOG_DIR') log_dir = os.path.join(temp_dir, 'weightslab_logs') os.makedirs(log_dir, exist_ok=True) @@ -155,46 +155,46 @@ def set_log_directory(new_log_dir): """ Updates the log file location to a new directory. Moves the existing log file from temp location to the new directory. - + This is automatically called when root_log_dir is resolved in training scripts. Can also be called manually if you want to relocate logs. - + Args: new_log_dir (str): The new directory where logs should be saved. - + Example: >>> import weightslab as wl >>> # Logging starts in temp directory automatically >>> # Later, when you define your experiment directory: >>> wl.set_log_directory("./my_experiment/logs") >>> # Log file is moved from temp to ./my_experiment/logs/ - + Note: - The log file keeps its original timestamped filename - All subsequent logs are written to the new location - The old temp directory log is moved (not copied) """ global _TMP_DIR_PATH, _LOG_FILE_PATH, _FILE_HANDLER - + if not _LOG_FILE_PATH or not _FILE_HANDLER: logging.warning("No log file to relocate. Call setup_logging() first.") return - + # Create new log directory os.makedirs(new_log_dir, exist_ok=True) - + # Generate new log file path with same filename old_filename = os.path.basename(_LOG_FILE_PATH) new_log_path = os.path.join(new_log_dir, old_filename) - + # Get root logger root_logger = logging.getLogger() - + # Flush and close current file handler _FILE_HANDLER.flush() _FILE_HANDLER.close() root_logger.removeHandler(_FILE_HANDLER) - + # Move the log file to new location try: if os.path.exists(_LOG_FILE_PATH): @@ -202,18 +202,18 @@ def set_log_directory(new_log_dir): logging.info(f"Log file moved from {_LOG_FILE_PATH} to {new_log_path}") except Exception as e: logging.warning(f"Could not move log file: {e}. Creating new log file at {new_log_path}") - + # Update global path _LOG_FILE_PATH = new_log_path _TMP_DIR_PATH = new_log_dir - + # Create new file handler at new location formatter = logging.Formatter(FORMAT, datefmt=DATE_FORMAT) _FILE_HANDLER = logging.FileHandler(_LOG_FILE_PATH, mode='a', encoding='utf-8') _FILE_HANDLER.setLevel(logging.DEBUG) _FILE_HANDLER.setFormatter(formatter) root_logger.addHandler(_FILE_HANDLER) - + logging.info(f"Log directory updated to: {new_log_dir}") logging.info(f"Log file: {_LOG_FILE_PATH}") @@ -271,7 +271,7 @@ def print(first_element, *other_elements, sep=' ', **kwargs): new_log_dir = os.path.join(tempfile.gettempdir(), 'weightslab_test_logs') print(f'Relocating logs to: {new_log_dir}') set_log_directory(new_log_dir) - + # Test 5: Log after relocation print('This is a message after log relocation.', 'All good.') print(f'New log file location: {_LOG_FILE_PATH}') From 0b135244e2b907aba7c9fa6f89a4d9aeb42faf5d Mon Sep 17 00:00:00 2001 From: GuillaumePELLUET Date: Sat, 26 Sep 2026 01:16:27 +0200 Subject: [PATCH 17/29] fix cls example --- .../PyTorch/wl-classification/config.yaml | 6 +- .../PyTorch/wl-classification/main.py | 738 ++++-------------- 2 files changed, 136 insertions(+), 608 deletions(-) diff --git a/weightslab/examples/PyTorch/wl-classification/config.yaml b/weightslab/examples/PyTorch/wl-classification/config.yaml index 35139e70..8882c477 100644 --- a/weightslab/examples/PyTorch/wl-classification/config.yaml +++ b/weightslab/examples/PyTorch/wl-classification/config.yaml @@ -1,7 +1,7 @@ # Global configuration experiment_name: mnist_classification device: cpu -training_steps_to_do: null # Set to null for infinite training until manually stopped - behavior set by the user in main script +# training_steps_to_do: null # Set to null for infinite training until manually stopped - behavior set by the user in main script # root_log_dir: # Path to save current experiment checkpoints and data checkpoint_manager: @@ -35,12 +35,12 @@ serving_grpc: true data: train_loader: shuffle: true - batch_size: 16 + batch_size: 64 # max_samples: 256 val_loader: shuffle: false # max_samples: 256 - batch_size: 64 + batch_size: 128 test_loader: shuffle: false # max_samples: 256 diff --git a/weightslab/examples/PyTorch/wl-classification/main.py b/weightslab/examples/PyTorch/wl-classification/main.py index 519975d1..29070508 100644 --- a/weightslab/examples/PyTorch/wl-classification/main.py +++ b/weightslab/examples/PyTorch/wl-classification/main.py @@ -1,7 +1,5 @@ import itertools -import math import os -import random import ssl import time import logging @@ -16,12 +14,8 @@ import yaml import tqdm -import json -import numpy as np -import pandas as pd import torch import torch.nn as nn -import torch.nn.functional as F import torch.optim as optim from torchvision import datasets, transforms @@ -37,6 +31,42 @@ logger = logging.getLogger(__name__) + +# ============================================================================= +# Model and recipe +# ============================================================================= +class CNN(nn.Module): + """2 conv + 2 fc layers (1.2M parameters), raw logits out, optional dropout. + + Raw logits, not a softmax: `nn.CrossEntropyLoss` applies its own log-softmax. + A softmax here would be applied twice, which still trains but squashes the + per-sample loss values -- and those values are what you sort and filter on + in the Studio. + """ + + def __init__(self, dropout=False): + super().__init__() + self.input_shape = (1, 1, 28, 28) + self.conv1 = nn.Conv2d(1, 32, 3) + self.relu1 = nn.ReLU() + self.conv2 = nn.Conv2d(32, 64, 3) + self.relu2 = nn.ReLU() + self.pool = nn.MaxPool2d(2) + self.drop1 = nn.Dropout(0.25 if dropout else 0.0) + self.flatten = nn.Flatten() + self.fc1 = nn.Linear(64 * 12 * 12, 128) + self.relu3 = nn.ReLU() + self.drop2 = nn.Dropout(0.5 if dropout else 0.0) + self.fc2 = nn.Linear(128, 10) + + def features(self, x): + x = self.drop1(self.pool(self.relu2(self.conv2(self.relu1(self.conv1(x)))))) + return self.relu3(self.fc1(self.flatten(x))) + + def forward(self, x): + return self.fc2(self.drop2(self.features(x))) + + # ============================================================================= # Custom MNIST Dataset with Filepath Metadata # ============================================================================= @@ -124,126 +154,17 @@ def __getitem__(self, idx): return image, idx, label -# ============================================================================= -# Model and recipe -# ============================================================================= -class CNN(nn.Module): - """2 conv + 2 fc layers (1.2M parameters), raw logits out, optional dropout. - - Raw logits, not a softmax: `nn.CrossEntropyLoss` applies its own log-softmax. - A softmax here would be applied twice, which still trains but squashes the - per-sample loss values -- and those values are what you sort and filter on - in the Studio. - """ - - def __init__(self, dropout=False): - super().__init__() - self.input_shape = (1, 1, 28, 28) - self.conv1 = nn.Conv2d(1, 32, 3) - self.relu1 = nn.ReLU() - self.conv2 = nn.Conv2d(32, 64, 3) - self.relu2 = nn.ReLU() - self.pool = nn.MaxPool2d(2) - self.drop1 = nn.Dropout(0.25 if dropout else 0.0) - self.flatten = nn.Flatten() - self.fc1 = nn.Linear(64 * 12 * 12, 128) - self.relu3 = nn.ReLU() - self.drop2 = nn.Dropout(0.5 if dropout else 0.0) - self.fc2 = nn.Linear(128, 10) - - def features(self, x): - x = self.drop1(self.pool(self.relu2(self.conv2(self.relu1(self.conv1(x)))))) - return self.relu3(self.fc1(self.flatten(x))) - - def forward(self, x): - return self.fc2(self.drop2(self.features(x))) - - -def gpu_augment(x, rot_deg=10.0, shift=0.1, scale=(0.9, 1.1)): - """Random affine per image (rotation, shift, scale), on the batch, on the GPU.""" - b = x.size(0) - ang = (torch.rand(b, device=x.device) * 2 - 1) * rot_deg * math.pi / 180 - s = torch.empty(b, device=x.device).uniform_(*scale) - tx = (torch.rand(b, device=x.device) * 2 - 1) * shift * 2 - ty = (torch.rand(b, device=x.device) * 2 - 1) * shift * 2 - cos, sin = torch.cos(ang) / s, torch.sin(ang) / s - theta = torch.stack([torch.stack([cos, -sin, tx], 1), torch.stack([sin, cos, ty], 1)], 1) - grid = F.affine_grid(theta, x.shape, align_corners=False) - return F.grid_sample(x, grid, align_corners=False, padding_mode="zeros") - - -def lr_at(step, total_steps, base_lr, sched): - """const, or warmcos: linear warmup over 15% of the steps, then cosine decay to ~0.""" - if sched == "const": - return base_lr - warm = max(1, int(0.15 * total_steps)) - if step <= warm: - return base_lr * (0.04 + 0.96 * step / warm) - p = (step - warm) / max(1, total_steps - warm) - return 1e-6 + base_lr * 0.5 * (1 + math.cos(math.pi * p)) - - -def get_val_ids(targets, path): - """Fixed validation split: 500 train images per class, never trained on. - - Fixed RNG seed, so the split is identical on every machine and every rerun. - Delete the file to draw a new one. - """ - if path and os.path.exists(path): - return sorted(int(l) for l in open(path) if l.strip()) - y = targets.numpy() if hasattr(targets, "numpy") else np.asarray(targets) - rng = np.random.RandomState(2026) - ids = [] - for c in range(10): - ids += rng.choice(np.where(y == c)[0], 500, replace=False).tolist() - ids = sorted(int(i) for i in ids) - if path: - open(path, "w").write("\n".join(map(str, ids)) + "\n") - return ids - - -class SubsetWithIds(Dataset): - """A view over `base` restricted to `ids`, in the order given. - - WeightsLab numbers each split's samples contiguously, so a sample id is a - position within the split, not an MNIST index. `ids` is the lookup back: - the n-th sample of this loader is `ids[n]` in MNIST, which is also the n-th - line of the matching train_ids.txt / val_ids.txt. - """ - - def __init__(self, base, ids): - self.base = base - self.ids = list(ids) - - def __len__(self): - return len(self.ids) - - def __getitem__(self, i): - return self.base[self.ids[i]] - - # ----------------------------------------------------------------------------- # Train / Test functions # ----------------------------------------------------------------------------- -def train(loader, model, optimizer, criterion_mlt, device, step, recipe): - """Single training step using the tracked dataloader + watched loss. - - `recipe` carries the run's knobs: augmentation and the LR schedule. - """ +def train(loader, model, optimizer, criterion_mlt, device): + """Single training step using the tracked dataloader + watched loss.""" with wl.guard_training_context: (inputs, ids, labels) = next(loader) inputs = inputs.to(device) labels = labels.to(device) - # Augment on the GPU, on the batch - if recipe["augment"]: - inputs = gpu_augment(inputs) - - # LR schedule is applied per step, from the model's age - for g in optimizer.param_groups: - g["lr"] = lr_at(step, recipe["total_steps"], recipe["lr"], recipe["schedule"]) - # Infer optimizer.zero_grad() preds_raw = model(inputs) @@ -254,9 +175,7 @@ def train(loader, model, optimizer, criterion_mlt, device, step, recipe): else: preds = preds_raw.argmax(dim=1, keepdim=True) - # Loss is a watched object => pass metadata for logging/stats. - # per_sample=True on the criterion keeps one loss value per digit per - # visit, so each sample has a loss history you can inspect in the Studio. + # Loss is a watched object => pass metadata for logging/stats loss_batch_mlt = criterion_mlt( preds_raw.float(), labels.long(), @@ -272,16 +191,9 @@ def train(loader, model, optimizer, criterion_mlt, device, step, recipe): return total_loss.detach().cpu().item() -def test(loader, model, criterion_mlt, metric_mlt, device, test_loader_len, split="test"): - """Full evaluation pass over one evaluation loader. - - The metric is reset first: torchmetrics accumulates across `update` calls, - so without this each evaluation would report a running average over every - previous evaluation instead of the current one -- and the val and test - loaders would pollute each other's number. - """ +def test(loader, model, criterion_mlt, metric_mlt, device, test_loader_len): + """Full evaluation pass over the test loader.""" losses = torch.tensor(0.0, device=device) - metric_mlt.reset() for (inputs, ids, labels) in loader: with wl.guard_testing_context: @@ -314,8 +226,8 @@ def test(loader, model, criterion_mlt, metric_mlt, device, test_loader_len, spli # Log per-sample metric alongside signals; persists via the storer signals = { - f"{split}_metric/Accuracy_per_sample": acc_per_sample, - f"{split}_metric/Inverse_Accuracy_per_sample": acc_reversed_per_sample, + "test_metric/Accuracy_per_sample": acc_per_sample, + "test_metric/Inverse_Accuracy_per_sample": acc_reversed_per_sample, } wl.save_signals( preds_raw=outputs, @@ -331,298 +243,6 @@ def test(loader, model, criterion_mlt, metric_mlt, device, test_loader_len, spli return loss.detach().cpu().item(), metric.detach().cpu().item() -def _set_dropout(model, enabled: bool) -> None: - """Toggle dropout between phases without rebuilding the model. - - Rebuilding would create a NEW watched model (new hash, new checkpoint - lineage) and lose the step-0 weights we are about to restore, so the rates - are set in place on the existing modules instead. - """ - for name, p in (("drop1", 0.25), ("drop2", 0.5)): - mod = getattr(model, name, None) - if mod is not None: - mod.p = p if enabled else 0.0 - - -def _reload_initial_weights(wl, model, optimizer, init_state, log_dir) -> dict: - """Put the model back to its step-0 random initialisation. - - Prefers WeightsLab's own step-0 checkpoint (what "reload at model age 0" - means from the Studio); falls back to the in-memory snapshot taken before - phase A. Either way the result is verified against that snapshot, so the - number reported is a measurement and not a claim. - """ - source = "in-memory snapshot" - try: - from weightslab.backend.ledgers import get_checkpoint_manager - cm = get_checkpoint_manager() - exp_hash = getattr(cm, "current_exp_hash", None) or cm.get_latest_hash() - cm.load_checkpoint(exp_hash=exp_hash, load_model=False, load_weights=True, - load_config=False, load_data=False, target_step=0, force=True) - source = "step-0 checkpoint" - except Exception as exc: - print(f"[reset] step-0 checkpoint reload unavailable ({exc}); using the snapshot", flush=True) - - def _drift(): - cur = model.state_dict() - return max((cur[k].float() - v.float()).abs().max().item() - for k, v in init_state.items() if k in cur and torch.is_tensor(cur[k])) - - worst = _drift() - if worst > 0: - model.load_state_dict(init_state, strict=False) # snapshot is authoritative - after = _drift() - source += f" + snapshot (checkpoint differed by {worst:.3g})" - worst = after - - # Adam moments must go too: keeping them would carry phase A's gradient - # history into a run that is supposed to start from scratch. - optimizer.state = type(optimizer.state)() - return {"source": source, "max_abs_diff": worst} - - -def _fmt_si(n: float, unit: str) -> str: - """1199882 -> '1.2M'. Model cards quote these rounded, so match that.""" - for div, suf in ((1e9, "G"), (1e6, "M"), (1e3, "K")): - if n >= div: - return f"{n / div:.1f}{suf}{unit}" - return f"{n:.0f}{unit}" - - -def _card(title, acc, params, n_data) -> None: - """The four-field summary, one per experiment. - - Parameters and FLOPs are identical for both runs by construction -- same - architecture, same input size. That is the point: only NumberTrainingData - moves, so any accuracy difference is attributable to the data, not capacity. - """ - print(f" {title}") - print(f" Accuracy: {acc:.2f}%") - print(f" Parameters: {_fmt_si(params, '')}") - print(f" NumberTrainingData: {n_data:,}") - - - -def _summarise(results, n_train, goldset_size) -> None: - """Print the comparison the experiment exists to produce.""" - a, b = results["signal"], results["goldset"] - macs = results["flops_fwd_per_sample"] - params = results["model_params"] - fpv = macs * 3 # fwd + bwd, per sample-visit, in MACs - - print("\n" + "=" * 74) - print(" MODEL CARDS") - print("=" * 74) - _card("full trainset (phase A: signal run)", - a["best"].get("test_acc", float("nan")), params, n_train) - print() - _card("goldset (phase B: retrained from step-0 weights)", - b["best"].get("test_acc", float("nan")), params, goldset_size) - - print("\n" + "=" * 74) - print(" TRAINING COST (what actually differs)") - print("=" * 74) - print(f"{'':24s}{'full trainset':>18s}{'goldset':>18s}") - rows = [ - ("training digits", f"{n_train:,}", f"{goldset_size:,}"), - ("steps x batch", f"{a['steps']:,} x {a['batch_size']}", f"{b['steps']:,} x {b['batch_size']}"), - ("sample visits", f"{a['sample_visits']:,}", f"{b['sample_visits']:,}"), - ("epochs over own set", f"{a['sample_visits'] / max(1, n_train):.1f}", - f"{b['sample_visits'] / max(1, goldset_size):.1f}"), - ("train TFLOPs", f"{a['sample_visits'] * fpv * 2 / 1e12:.2f}", - f"{b['sample_visits'] * fpv * 2 / 1e12:.2f}"), - ("TFLOPs to best val", f"{a['best'].get('tflops', 0) * 2:.2f}", - f"{b['best'].get('tflops', 0) * 2:.2f}"), - ("train seconds", f"{a['train_seconds']:.1f}", f"{b['train_seconds']:.1f}"), - ("eval seconds", f"{a['eval_seconds']:.1f}", f"{b['eval_seconds']:.1f}"), - ("wall seconds", f"{a['wall_seconds']:.1f}", f"{b['wall_seconds']:.1f}"), - ("best val acc %", f"{a['best'].get('val_acc', float('nan')):.2f}", - f"{b['best'].get('val_acc', float('nan')):.2f}"), - ("test @ best val %", f"{a['best'].get('test_acc', float('nan')):.2f}", - f"{b['best'].get('test_acc', float('nan')):.2f}"), - ("step of best val", f"{a['best'].get('step', 0):,}", f"{b['best'].get('step', 0):,}"), - ] - for label, av, bv in rows: - print(f"{label:24s}{av:>18s}{bv:>18s}") - - ep_a = n_train * fpv * 2 / 1e12 - ep_b = goldset_size * fpv * 2 / 1e12 - saved = f"{(1 - ep_b / ep_a) * 100:.1f}% less per epoch" if ep_a > 0 else "n/a" - print("\n one epoch over its own training set:") - print(f" full trainset: {ep_a:.3f} TFLOPs over {n_train:,} digits") - print(f" goldset: {ep_b:.3f} TFLOPs over {goldset_size:,} digits ({saved})") - print(" Same architecture, same per-sample FLOPs -- training cost scales purely") - print(" with sample-visits, so the saving is in the data, not the model.") - print("=" * 74) - - -# ============================================================================ -# Experiment phases: signal run -> goldset -> retrain from the initial weights -# ============================================================================ -def model_flops_per_sample(model, device) -> int: - """Forward MACs for one sample, counted from the conv/linear layers. - - The model is identical in both phases, so this is a constant; what differs - between the phases is only how many sample-visits each one pays for. - """ - macs = {"n": 0} - hooks = [] - - def conv_hook(mod, inp, out): - macs["n"] += out.numel() * mod.in_channels // mod.groups * mod.kernel_size[0] * mod.kernel_size[1] - - def lin_hook(mod, inp, out): - macs["n"] += out.numel() * mod.in_features - - for m in model.modules(): - if isinstance(m, nn.Conv2d): - hooks.append(m.register_forward_hook(conv_hook)) - elif isinstance(m, nn.Linear): - hooks.append(m.register_forward_hook(lin_hook)) - was_training = model.training - model.eval() - with torch.no_grad(): - model(torch.zeros(1, 1, 28, 28, device=device)) - model.train(was_training) - for h in hooks: - h.remove() - if macs["n"] == 0: - raise RuntimeError( - "FLOPs count came back 0 -- no Conv2d/Linear was visited. Pass the " - "UNWRAPPED model (the object built before wl.watch_or_edit): a watched " - "model is a ModelInterface whose .modules() yields only itself.") - return int(macs["n"]) - - -def build_goldset(wl, steps_per_pass, labels_by_id, pool_ids, gcfg): - """Per class, the digits with the highest mean loss over passes 2..N, minus - the noisy ones (median loss over all passes above the threshold). - - Reads the per-sample loss history WeightsLab recorded during the signal run - -- no second inference pass, no extra bookkeeping of our own. - """ - wl.drain_signals() - rows = wl.query_signal_history("train-loss-CE") - hist = pd.DataFrame(rows, columns=["sample_id", "step", "loss", "run"][:len(rows[0])]) if rows else pd.DataFrame() - if hist.empty: - raise SystemExit("[goldset] no train-loss-CE history recorded; cannot build the goldset") - hist["sample_id"] = hist["sample_id"].astype(int) - hist["loss"] = hist["loss"].astype(float) - # "Pass n" is the n-th time THIS digit was seen, taken from its own visit - # order rather than from step arithmetic: the recorded step numbering is not - # guaranteed to start at 1, and a uniform shuffle is what makes the two - # equivalent in the first place. - hist = hist.sort_values(["sample_id", "step"], kind="stable") - hist["pass"] = hist.groupby("sample_id").cumcount() - per_pass = hist.groupby(["sample_id", "pass"])["loss"].last().unstack() - - noisy = per_pass.median(axis=1, skipna=True) > float(gcfg.get("noisy_median_loss", 1.0)) - first = 1 if gcfg.get("skip_first_pass", True) else 0 - n_passes = int(per_pass.shape[1]) - if n_passes <= first: - raise SystemExit( - f"[goldset] only {n_passes} pass(es) recorded; the rule scores passes " - f"{first + 1}..N, so it needs at least {first + 2}. Per-sample loss over " - f"fewer passes measures batch placement, not difficulty.") - score = per_pass.iloc[:, first:].mean(axis=1, skipna=True) - - pool = set(int(i) for i in pool_ids) - cand = pd.DataFrame({"score": score, "noisy": noisy.reindex(score.index, fill_value=False)}) - cand = cand[cand.index.map(lambda i: int(i) in pool) & ~cand["noisy"] & cand["score"].notna()] - cand["label"] = [labels_by_id[int(i)] for i in cand.index] - - k = int(gcfg.get("per_class", 360)) - gold, per_class_got = [], {} - for c in sorted(cand["label"].unique()): - g = cand[cand["label"] == c].sort_values("score", ascending=False, kind="stable") - picked = [int(i) for i in g.head(k).index] - gold += picked - per_class_got[int(c)] = len(picked) - - # A goldset that is empty, or short of its per-class quota, means the rule or - # the history is wrong -- stop rather than retrain on a malformed subset. - if not gold: - raise SystemExit("[goldset] selection is empty; refusing to continue") - short = {c: n for c, n in per_class_got.items() if n < k} - if short: - raise SystemExit(f"[goldset] classes below the {k}/class quota: {short}") - if len(set(gold)) != len(gold): - raise SystemExit("[goldset] duplicate ids in the selection") - - return sorted(gold), {"noisy": int(noisy.sum()), "scored": int(len(cand)), - "passes": int(per_pass.shape[1]), "per_class": per_class_got} - - -def run_phase(name, total_steps, recipe, eval_every, ctx): - """Train one phase, evaluating validation on a cadence and TEST ONLY when - validation reaches a new maximum. - - Test accuracy is therefore never used to steer anything -- it is read at - exactly the points where validation says the model just got better, which - is the only honest moment to look at it. - """ - wl, model, optimizer = ctx["wl"], ctx["model"], ctx["optimizer"] - train_loader, val_loader, test_loader = ctx["train_loader"], ctx["val_loader"], ctx["test_loader"] - device = ctx["device"] - - flops_per_visit = ctx["flops_fwd_per_sample"] * 3 # fwd + bwd ~= 3x fwd - bs = recipe["batch_size"] - - best_val, best = -1.0, {} - train_seconds = eval_seconds = 0.0 - visits = 0 - history = [] - t_phase = time.perf_counter() - - for step in range(1, total_steps + 1): - t0 = time.perf_counter() - loss = train(train_loader, model, optimizer, ctx["train_criterion"], device, step, recipe) - train_seconds += time.perf_counter() - t0 - visits += bs - - if step % eval_every == 0 or step == total_steps: - t0 = time.perf_counter() - val_loss, val_acc = test(val_loader, model, ctx["val_criterion"], ctx["val_metric"], - device, ctx["val_loader_len"], split="val") - row = {"phase": name, "step": step, "train_loss": round(loss, 5), - "val_acc": round(val_acc, 4), "test_acc": None, - "visits": visits, "tflops": round(visits * flops_per_visit / 1e12, 4), - "train_s": round(train_seconds, 1)} - signals = {"val/accuracy": val_acc} - - if val_acc > best_val: # new best validation -> read test once - best_val = val_acc - _, test_acc = test(test_loader, model, ctx["test_criterion"], ctx["test_metric"], - device, ctx["test_loader_len"], split="test") - row["test_acc"] = round(test_acc, 4) - signals["test/accuracy"] = test_acc - best = {"step": step, "val_acc": val_acc, "test_acc": test_acc, - "visits": visits, "tflops": row["tflops"], - "train_s": round(train_seconds, 1)} - eval_seconds += time.perf_counter() - t0 - - wl.save_model_signals(signals) - history.append(row) - mark = " <- new best val, test read" if row["test_acc"] is not None else "" - print(f"[{name}] step {step:5d}/{total_steps} loss {loss:.4f} " - f"val {val_acc:6.2f}% test {row['test_acc'] if row['test_acc'] is not None else ' - '}" - f" {row['tflops']:.2f} TFLOPs{mark}", flush=True) - - return { - "phase": name, - "steps": total_steps, - "batch_size": bs, - "sample_visits": visits, - "train_seconds": round(train_seconds, 1), - "eval_seconds": round(eval_seconds, 1), - "wall_seconds": round(time.perf_counter() - t_phase, 1), - "flops_fwd_per_sample": ctx["flops_fwd_per_sample"], - "train_tflops": round(visits * flops_per_visit / 1e12, 4), - "best": best, - "history": history, - } - - # ----------------------------------------------------------------------------- # Main # ----------------------------------------------------------------------------- @@ -642,20 +262,6 @@ def run_phase(name, total_steps, recipe, eval_every, ctx): parameters.setdefault("device", "auto") parameters.setdefault("training_steps_to_do", 1000000) parameters.setdefault("eval_full_to_steps_ratio", 50) - parameters.setdefault("seed", 0) - parameters.setdefault("augment", False) - parameters.setdefault("dropout", False) - parameters.setdefault("optimizer", {}).setdefault("schedule", "const") - parameters["optimizer"].setdefault("weight_decay", 0.0) - - # Deterministic seeding: a rerun reproduces the same shuffle, and therefore - # the same per-digit visit order. - seed = int(parameters["seed"]) - random.seed(seed) - np.random.seed(seed) - torch.manual_seed(seed) - torch.backends.cudnn.deterministic = True - torch.backends.cudnn.benchmark = False # Experiment name exp_name = parameters["experiment_name"] @@ -688,14 +294,15 @@ def run_phase(name, total_steps, recipe, eval_every, ctx): enable_h5_persistence = parameters.get("enable_h5_persistence", True) # Model - _model = CNN(dropout=bool(parameters["dropout"])).to(device) - model = wl.watch_or_edit(_model, flag="model", device=device) + _model = CNN().to(device) + model = wl.watch_or_edit(_model, flag="model", device=device, + compute_dependencies=True, + forced_model_wrapping=True, + skip_previous_auto_load=True) # Optimizer - opt_cfg = parameters["optimizer"] - lr = opt_cfg.get("lr", 0.001) - _optimizer = optim.Adam( - model.parameters(), lr=lr, weight_decay=float(opt_cfg["weight_decay"])) + lr = parameters.get("optimizer", {}).get("lr", 0.01) + _optimizer = optim.Adam(model.parameters(), lr=lr) optimizer = wl.watch_or_edit( _optimizer, flag="optimizer", @@ -717,10 +324,9 @@ def run_phase(name, total_steps, recipe, eval_every, ctx): # Read data config for all loaders train_cfg = parameters.get("data", {}).get("train_loader", {}) - val_cfg = parameters.get("data", {}).get("val_loader", {}) test_cfg = parameters.get("data", {}).get("test_loader", {}) - _full_train = MNISTCustomDataset( + _train_dataset = MNISTCustomDataset( root=data_root, train=True, download=should_download, @@ -743,35 +349,7 @@ def run_phase(name, total_steps, recipe, eval_every, ctx): max_samples=test_cfg.get("max_samples", None) ) - # Split the 60,000 MNIST train images into the training pool and a held-out - # validation set BEFORE anything is registered: the two loaders then own - # disjoint images, so nothing has to be discarded afterwards and no image is - # registered twice. - # - # WeightsLab numbers each split's samples contiguously from where the - # previous one ended, so a sample id here is a position in the split, not an - # MNIST index. Both index lists are written next to this config, in loader - # order, so the mapping back is a lookup: - # - # mnist_index = train_ids[wl_sample_id] (train split) - # mnist_index = val_ids[wl_sample_id - len(train_ids)] (val split) - val_ids_path = parameters.get("val_ids") - if val_ids_path and not os.path.isabs(val_ids_path): - val_ids_path = os.path.join(os.path.dirname(__file__), val_ids_path) - val_ids = get_val_ids(_full_train.mnist.targets, val_ids_path) - held_out = set(val_ids) - train_ids = [i for i in range(len(_full_train)) if i not in held_out] - assert not (set(train_ids) & held_out), "train ids overlap the validation split" - - if val_ids_path: # same order the loader serves them in - with open(os.path.join(os.path.dirname(val_ids_path), "train_ids.txt"), "w") as fh: - fh.write("\n".join(map(str, train_ids)) + "\n") - - _train_dataset = SubsetWithIds(_full_train, train_ids) - _val_dataset = SubsetWithIds(_full_train, val_ids) - n_train = len(_train_dataset) - - # Create tracked loaders for train, val and test + # Create tracked loaders for train, test, and test train_loader = wl.watch_or_edit( _train_dataset, flag="data", @@ -784,18 +362,6 @@ def run_phase(name, total_steps, recipe, eval_every, ctx): preload_metadata=False, enable_h5_persistence=enable_h5_persistence ) - val_loader = wl.watch_or_edit( - _val_dataset, - flag="data", - loader_name="val_loader", - batch_size=val_cfg.get("batch_size", 500), - shuffle=val_cfg.get("shuffle", False), - is_training=False, - compute_hash=False, - preload_labels=True, - preload_metadata=False, - enable_h5_persistence=enable_h5_persistence - ) test_loader = wl.watch_or_edit( _test_dataset, flag="data", @@ -809,142 +375,104 @@ def run_phase(name, total_steps, recipe, eval_every, ctx): enable_h5_persistence=enable_h5_persistence ) - # 8 passes over 55,000 digits at batch 64 = 8 x 860 = 6,880 steps. - batch_size = train_cfg.get("batch_size", 16) - steps_per_pass = -(-n_train // batch_size) - if parameters.get("epochs"): - parameters["training_steps_to_do"] = int(parameters["epochs"]) * steps_per_pass - total_steps = int(parameters["training_steps_to_do"]) - recipe = { - "augment": bool(parameters["augment"]), - "lr": float(lr), - "schedule": opt_cfg["schedule"], - "total_steps": total_steps, - } - # Losses & metrics (watched objects – they log themselves) train_criterion = wl.watch_or_edit( nn.CrossEntropyLoss(reduction="none"), - flag="loss", signal_name="train-loss-CE", per_sample=True, log=True) + flag="loss", signal_name="train-loss-CE", log=True) test_criterion = wl.watch_or_edit( nn.CrossEntropyLoss(reduction="none"), flag="loss", signal_name="test-loss-CE", log=True) - val_criterion = wl.watch_or_edit( - nn.CrossEntropyLoss(reduction="none"), - flag="loss", signal_name="val-loss-CE", log=True) - metric = wl.watch_or_edit( Accuracy(task="multiclass", num_classes=10).to(device), flag="metric", signal_name="metric-ACC", log=True) - val_metric = wl.watch_or_edit( - Accuracy(task="multiclass", num_classes=10).to(device), - flag="metric", signal_name="metric-ACC-val", log=True) # Start WeightsLab services (gRPC only, no CLI) wl.serve( serving_grpc=parameters.get("serving_grpc", False) ) - # ---- constants shared by both phases ------------------------------------- - val_loader_len = len(val_loader) - test_loader_len = len(test_loader) - flops_fwd = model_flops_per_sample(_model, device) # unwrapped: see the docstring - labels_by_id = {i: int(_full_train.mnist.targets[m]) for i, m in enumerate(train_ids)} - gcfg = parameters.get("goldset", {}) or {} - - ctx = dict(wl=wl, model=model, optimizer=optimizer, device=device, - train_loader=train_loader, val_loader=val_loader, test_loader=test_loader, - train_criterion=train_criterion, val_criterion=val_criterion, - test_criterion=test_criterion, val_metric=val_metric, test_metric=metric, - val_loader_len=val_loader_len, test_loader_len=test_loader_len, - flops_fwd_per_sample=flops_fwd) - - p_sig = parameters["phases"]["signal"] - p_gold = parameters["phases"]["goldset"] - sig_steps = int(p_sig["epochs"]) * (-(-n_train // int(p_sig["batch_size"]))) - - print("=" * 72) - print(" TWO-PHASE EXPERIMENT (one process, one experiment directory)") - print(f" splits: train={n_train} val={len(_val_dataset)} test={len(_test_dataset)}") - print(f" model: {sum(p.numel() for p in model.parameters()):,} params, " - f"{flops_fwd / 1e6:.2f} MFLOPs forward per sample ({flops_fwd * 3 / 1e6:.2f} incl. backward)") - print(f" phase A (signal): {sig_steps} steps x batch {p_sig['batch_size']} " - f"= {p_sig['epochs']} passes over {n_train} digits") - print(f" phase B (goldset): {p_gold['training_steps_to_do']} steps x batch {p_gold['batch_size']}") - print(f" logs: {log_dir}") - print("=" * 72 + "\n") - - # Hand control to WeightsLab before the first guarded step: without this the - # training guard never opens and the loop stalls with the process alive. - wl.start_training(timeout=int(parameters.get("start_training_timeout", 1))) - - results = {"splits": {"train": n_train, "val": len(_val_dataset), "test": len(_test_dataset)}, - "model_params": int(sum(p.numel() for p in model.parameters())), - "flops_fwd_per_sample": flops_fwd} - t_all = time.perf_counter() - - # ---- the initial random weights, kept so phase B can start from them ------ - # Tensors only: a watched model's state_dict also carries scalar bookkeeping - # (the age counter), which has nothing to clone or compare. - init_state = {k: v.detach().clone() for k, v in model.state_dict().items() - if torch.is_tensor(v)} - - # ---- PHASE A: signal run over the whole training pool --------------------- - recipe_sig = {"augment": bool(p_sig["augment"]), "lr": float(p_sig["lr"]), - "schedule": p_sig["schedule"], "total_steps": sig_steps, - "batch_size": int(p_sig["batch_size"])} - train_loader.set_batch_size(int(p_sig["batch_size"])) - _set_dropout(model, bool(p_sig["dropout"])) - for g in optimizer.param_groups: - g["lr"] = float(p_sig["lr"]) - results["signal"] = run_phase("signal", sig_steps, recipe_sig, int(p_sig["eval_every"]), ctx) - - # ---- build the goldset from what phase A recorded ------------------------- - t0 = time.perf_counter() - steps_per_pass = -(-n_train // int(p_sig["batch_size"])) - goldset, ginfo = build_goldset(wl, steps_per_pass, labels_by_id, range(n_train), gcfg) - wl.tag_samples(goldset, "goldset") - results["goldset_build"] = {**ginfo, "size": len(goldset), - "seconds": round(time.perf_counter() - t0, 1)} - mnist_ids = sorted(int(train_ids[i]) for i in goldset) - with open(os.path.join(log_dir, "goldset_ids_mnist.txt"), "w") as fh: - fh.write("\n".join(map(str, mnist_ids)) + "\n") - print(f"\n[goldset] {len(goldset)} digits ({len(goldset) / n_train:.2%} of the pool) " - f"from {ginfo['passes']} passes; {ginfo['noisy']} noisy excluded; " - f"tagged 'goldset'; MNIST ids -> goldset_ids_mnist.txt\n", flush=True) - - # ---- reload the INITIAL weights (model age 0), then train on the goldset --- - reload_info = _reload_initial_weights(wl, model, optimizer, init_state, log_dir) - results["reload"] = reload_info - print(f"[reset] restored the step-0 weights ({reload_info['source']}); " - f"max|w - w0| = {reload_info['max_abs_diff']}\n", flush=True) - - keep = set(int(i) for i in goldset) - wl.discard_samples([i for i in range(n_train) if i not in keep]) - - recipe_gold = {"augment": bool(p_gold["augment"]), "lr": float(p_gold["lr"]), - "schedule": p_gold["schedule"], "total_steps": int(p_gold["training_steps_to_do"]), - "batch_size": int(p_gold["batch_size"])} - train_loader.set_batch_size(int(p_gold["batch_size"])) - _set_dropout(model, bool(p_gold["dropout"])) - results["goldset"] = run_phase("goldset", int(p_gold["training_steps_to_do"]), - recipe_gold, int(p_gold["eval_every"]), ctx) - - # ---- report --------------------------------------------------------------- - results["wall_seconds_total"] = round(time.perf_counter() - t_all, 1) - _summarise(results, n_train, len(goldset)) - with open(os.path.join(log_dir, "experiment_results.json"), "w") as fh: - json.dump(results, fh, indent=1) - print(f"\n results -> {os.path.join(log_dir, 'experiment_results.json')}") + print("=" * 60) + print(" STARTING TRAINING") + print(f" Evaluation every {eval_full_to_train_steps_ratio} steps") + print(f" Dataset splits: train={len(_train_dataset)}, test={len(_test_dataset)}") + print(f" Logs will be saved to: {log_dir}") + print("=" * 60 + "\n") + + # Setup clean progress bar with custom format + # Training runs until YOU stop it -- from the studio's pause button, the CLI, + # or Ctrl+C. itertools.count() rather than range(training_steps_to_do): a + # predefined step budget ends the process mid-experiment, which is the + # opposite of how WeightsLab is used (inspect the curves, edit the data or + # the architecture, keep going). `training_steps_to_do` remains a live + # hyperparameter for the UI's own "run N more steps" control; it is not a + # ceiling on this loop. + if tqdm_display: + train_range = tqdm.tqdm( + itertools.count(), + desc="Training", + bar_format="{desc}: {n} steps [{elapsed}, {rate_fmt}] {bar} | {postfix}", + ncols=140, + position=0, + leave=True + ) + else: + train_range = itertools.count() + + # ============= + # Training Loop + wl.start_training(timeout=3) # Blocks and keeps the main thread alive while background services run. Optionally set a timeout (seconds) to auto-stop. + + train_loss = None + test_loss, test_metric = None, None + test_loader_len = len(test_loader) # Store length before wrapping with tqdm + for train_step in train_range: + age = model.get_age() if hasattr(model, "get_age") else train_step # Get model age in steps (not necessarily equal to train_step if model was reloaded or has seen more data than training steps) + + # Train one step + train_loss = train(train_loader, model, optimizer, train_criterion, device) + + # Periodic test evaluation + if age > 0 and age % eval_full_to_train_steps_ratio == 0: + # Test (no nested progress bar) + test_loss, test_metric = test( + test_loader, + model, + test_criterion, + metric, + device, + test_loader_len + ) + # Verbose + if verbose and not tqdm_display: + import sys + # Build compact progress message + msg = f"Step {train_step} (Age {age}): Loss={train_loss:.4f}" + if test_loss is not None: + msg += f" | Test={test_loss:.4f} ({test_metric:.1f}%)" + + # Clear line completely and print (pad to 100 chars to overwrite previous content) + sys.stdout.write(f"\r{msg:<100}") + sys.stdout.flush() + elif tqdm_display: + # Build compact postfix string + postfix_parts = [f"train_loss={train_loss:.4f}"] + if test_loss is not None: + postfix_parts.append(f"test_loss={test_loss:.4f}") + if test_metric is not None: + postfix_parts.append(f"test_acc={test_metric:.1f}%") + + train_range.set_postfix_str(" | ".join(postfix_parts)) + + print("\n" + "=" * 60) + print(f" Training completed in {time.time() - start_time:.2f} seconds") + print(f" Logs saved to: {log_dir}") + print("=" * 60) + + # Final export of signal history and data grid to root_log_dir wl.write_history() wl.write_dataframe() - # Keep the main thread alive so the Studio stays attached. Set - # keep_serving: false to exit once the results are written -- a batch sweep - # runs these back to back and must not block on the last one. - if parameters.get("keep_serving", True): - wl.keep_serving() - else: - print(" keep_serving: false -> exiting", flush=True) + # Keep the main thread alive to allow background serving threads to run + wl.keep_serving() From 21fcb4d055ddb2674f911ba9d2b81157e0d94c08 Mon Sep 17 00:00:00 2001 From: GuillaumePELLUET Date: Tue, 29 Sep 2026 20:44:17 +0200 Subject: [PATCH 18/29] se/start: OS-aware cert generation, certs dir fallback, CLI help + docs audit weightslab se - Windows now runs the PowerShell cert script (host openssl) first and only falls back to the bash script via WSL if it fails; --force-ubuntu forces the WSL bash script. A hung WSL used to block `se` forever with no output. - PowerShell now gets the native certs path (it used to get /mnt/c/..., which it resolves to C:\mnt\c\...). - With no CERTS_DIR argument, honour $WEIGHTSLAB_CERTS_DIR as documented. Certs directory resolution (UI `start --certs`, backend init, `se`) - $WEIGHTSLAB_CERTS_DIR first; ~/.weightslab-certs when the variable is unset, empty, not an absolute path (ignored with a warning), or holds no certs while ~/.weightslab-certs does. - _persist_certs_dir refuses non-absolute values, so a bogus value (e.g. a leaked mock repr) can no longer reach setx or ~/.bashrc. CLI help / docs audit - `start example --gen` pointed at wl-generation, renamed to wl-image-generation in #294; the command exited "example not found". - `start --port` help claimed a 50051 default; the real default is 8080. - The hidden `example` alias printed "==SUPPRESS==" in --help; `agent init`, `start --host` and `se CERTS_DIR` were missing from the help epilog. - docs: user_commands (5173 -> 8080, Envoy/Docker wording in tunnel, start/cli flags), agent.rst (/loop pipes lines into `weightslab cli`; there is no `weightslab pause`/`status`), wl-image-generation paths, certs lookup order. Tests - Sandbox WEIGHTSLAB_CERTS_DIR in the `se` tests (they leaked it into later tests and wrote a token into the real ~/.weightslab-certs), and pin GRPC_TLS_ENABLED=0 in the grpc_serve port tests (they failed on any machine with certs in ~/.weightslab-certs). - New coverage: cert script selection, --force-ubuntu, certs dir resolution, every example flag has a main.py, every subcommand has --help. Co-Authored-By: Claude Opus 5.5 --- AGENTS.md | 13 +- docs/agent.rst | 37 ++-- docs/examples/pytorch/generation.rst | 2 +- docs/usage/parameters.rst | 8 +- docs/user_commands.rst | 107 ++++++++--- docs/weights_studio/configuration.rst | 6 +- docs/weights_studio/security.rst | 17 +- docs/weights_studio_ui/more/configuration.rst | 6 +- docs/weights_studio_ui/more/security.rst | 17 +- tests/backend/test_cli.py | 179 +++++++++++++++++- tests/general/test_cli.py | 4 +- tests/test_secure_communication.py | 68 ++++++- .../services/test_trainer_services_server.py | 6 +- weightslab/AGENTS.md | 14 +- weightslab/cli.py | 176 ++++++++++------- weightslab/security/__init__.py | 4 +- weightslab/security/cert_auth_manager.py | 58 +++++- 17 files changed, 579 insertions(+), 143 deletions(-) diff --git a/AGENTS.md b/AGENTS.md index bf549c8c..5786cc5d 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -101,7 +101,9 @@ Working starting points live in UI deployment details (port, TLS, certs) are documented in `weightslab/docs/weights_studio.rst`. TLS is opt-in: run `weightslab se` once, -then `weightslab start --certs`. +then `weightslab start --certs`. On Windows `se` uses the PowerShell script and +the Windows `openssl`; `weightslab se --force-ubuntu` uses the bash script +through WSL instead. --- @@ -179,7 +181,7 @@ ones when debugging: | `GRPC_BACKEND_HOST` / `GRPC_BACKEND_PORT` | `0.0.0.0` / `50051` | Backend gRPC bind address. | | `GRPC_TLS_ENABLED` | `0` | TLS on the gRPC socket. Set `1` with `weightslab start --certs`. | | `GRPC_TLS_REQUIRE_CLIENT_AUTH` | `0` | mTLS. Must match what `weightslab start --certs` presents. | -| `WEIGHTSLAB_CERTS_DIR` | `~/.weightslab-certs` | Where cert files are looked up (single source of truth). | +| `WEIGHTSLAB_CERTS_DIR` | `~/.weightslab-certs` | Where cert files are looked up (single source of truth). Falls back to `~/.weightslab-certs` when unset, not an absolute path, or holding no certs. | | `GRPC_AUTH_TOKEN` | *(unset)* | Optional metadata-token auth on top of mTLS. | | `GRPC_MAX_MESSAGE_BYTES` | `268435456` (256 MB) | Raise it if large tensors/image batches fail. | | `WEIGHTSLAB_DISABLE_WATCHDOGS` | `0` | Set `1` when debugging with breakpoints (see §5). | @@ -220,6 +222,13 @@ can reach it on `:8080`; (3) **TLS mismatch** if using `--certs` — run `weightslab se` first and export `WEIGHTSLAB_CERTS_DIR`. For local debugging drop TLS entirely (omit `--certs`; `GRPC_TLS_ENABLED=0`). +**`weightslab se` hangs with no output (Windows).** +Only the WSL path can do this: `--force-ubuntu`, or the fallback after the +PowerShell script fails. There `bash` is the WSL launcher, the script's output +is captured, and there is no timeout, so a stuck WSL distro blocks forever. +Confirm with `wsl -e echo ok` (it hangs too). Fix with `wsl --shutdown`, or drop +`--force-ubuntu` so the PowerShell script runs. + **Changed an env var, restarted, but the UI still uses the old value.** - `VITE_*` is build-time → you must **rebuild** the frontend, not just restart. - `WS_*` / `BB_*` / `ENABLE_*` are injected at `weightslab start` time → you diff --git a/docs/agent.rst b/docs/agent.rst index 9519418d..7b6c7390 100644 --- a/docs/agent.rst +++ b/docs/agent.rst @@ -532,22 +532,27 @@ backend SDK agent directly. - **Syntax**: ``/loop m|h `` to start (minimum interval: 60s), ``/loop list`` to see running jobs, ``/loop stop `` to cancel one. -- **What it can do**: the loop's OpenCode session is told about the local - ``weightslab`` CLI, reachable over bash against the live training process: - - - ``weightslab pause`` / ``weightslab resume`` — freeze/resume weight updates - - ``weightslab discard `` — discard a sample by id - - ``weightslab agent query ""`` — hands the request to the - **backend SDK agent's** own intent pipeline, e.g. ``weightslab agent - query "discard samples where loss > 5 and tag them hard_examples"``. This - is how the loop reaches back into the database/history: it can't ask the - backend agent directly, but it can drive it through the CLI. - - ``weightslab status`` — a snapshot of hyperparameters/model/training state - - These four are what the loop's system prompt explicitly calls out, but bash - access means any other ``weightslab`` CLI verb is reachable too — e.g. - ``weightslab report`` to generate a narrative report for the loop to read - and act on. It may also read/edit training code directly and attempt to +- **What it can do**: the loop's OpenCode session controls the live training + process through the ``weightslab cli`` console. The console is interactive, + so the loop pipes in one line per bash call (EOF ends the session), e.g. + ``echo "status" | weightslab cli``. The lines its system prompt calls out: + + - ``pause`` / ``resume`` — freeze/resume weight updates + - ``discard `` — discard a sample by id + - ``agent query ""`` — hands the request to the + **backend SDK agent's** own intent pipeline, e.g. ``agent query "discard + samples where loss > 5 and tag them hard_examples"``. This is how the loop + reaches back into the database/history: it can't ask the backend agent + directly, but it can drive it through the console. The same verb builds + the experiment report (``agent query "Generate an experiment report."``); + later check-ins ask it to *update* that report rather than make a new one. + - ``status`` — component names and model age only (no metric or + hyperparameter values; ask ``agent query`` for those) + + These are console lines, not shell commands: there is no ``weightslab + status`` or ``weightslab pause``. Bash access means any other console verb + (see :doc:`weights_studio_cli/cli_console`) is reachable the same way. It + may also read/edit training code directly and attempt to restart a crashed process via bash — this is best-effort (no supervisor or PID handoff): it looks for the process, stops it if still running, and re-launches from whatever it can determine (shell history, a run script, diff --git a/docs/examples/pytorch/generation.rst b/docs/examples/pytorch/generation.rst index 7c283b50..f3a7e4c2 100644 --- a/docs/examples/pytorch/generation.rst +++ b/docs/examples/pytorch/generation.rst @@ -12,7 +12,7 @@ Generation / Anomaly Detection — MVTec (PyTorch) reconstruction
-**Example:** ``weightslab/examples/PyTorch/wl-generation/main.py`` +**Example:** ``weightslab/examples/PyTorch/wl-image-generation/main.py`` **Task:** Unsupervised anomaly detection on MVTec capsule images with a multi-task UNet (classification head + reconstruction head + contrastive loss). diff --git a/docs/usage/parameters.rst b/docs/usage/parameters.rst index 254fa0f9..5a4df2c3 100644 --- a/docs/usage/parameters.rst +++ b/docs/usage/parameters.rst @@ -339,9 +339,11 @@ Security & TLS - Default - Description * - ``WEIGHTSLAB_CERTS_DIR`` - - *(auto-generated)* - - Directory for TLS certificates and the gRPC auth token. - Auto-created under the user home dir when unset. + - ``~/.weightslab-certs`` + - Directory for TLS certificates and the gRPC auth token. Read first; + ``~/.weightslab-certs`` is used instead when it is unset, is not an + absolute path (ignored with a warning), or holds no certs while + ``~/.weightslab-certs`` does. * - ``GRPC_TLS_ENABLED`` - ``true`` - Enable TLS for the gRPC backend server. Set to ``false`` for diff --git a/docs/user_commands.rst b/docs/user_commands.rst index e885274b..ac516702 100644 --- a/docs/user_commands.rst +++ b/docs/user_commands.rst @@ -41,7 +41,7 @@ weightslab se .. code-block:: bash - weightslab se [certs_dir] [--force-certs] + weightslab se [certs_dir] [--force-certs] [--force-ubuntu] Generates TLS certificates and a gRPC auth token into a certs directory, then tells you to export ``WEIGHTSLAB_CERTS_DIR`` — the **single source of @@ -49,6 +49,32 @@ truth** the training backend, ``weightslab start --certs``, and any new shell all read to decide whether TLS/auth is on (derived purely from whether cert files exist in that directory). +The certificates come from a bundled script that needs ``openssl`` on +``PATH``. Which script runs depends on the OS: + +- **Linux / macOS:** the bash script (``generate-certs-auth-token.sh``). +- **Windows:** the PowerShell script (``generate-certs-auth-token.ps1``), using + the Windows ``openssl``. It also adds the dev CA to your user's trusted root + certificates, and Windows asks you to confirm. If the script fails, + ``weightslab se`` falls back to the bash script through WSL. + +Options: + +- ``certs_dir`` — directory for the certs and token (default: + ``$WEIGHTSLAB_CERTS_DIR``, else ``~/.weightslab-certs``). A + ``WEIGHTSLAB_CERTS_DIR`` that isn't an absolute path is ignored with a + warning. +- ``--force-certs`` — regenerate the certificates even if they already exist. +- ``--force-ubuntu`` — Windows only. Skip PowerShell and run the bash script + in your default WSL distribution (for example Ubuntu; ``wsl -l -v`` shows + which one is the default), with no fallback. Use it when you want the WSL + ``openssl``. This path does not add the CA to the Windows trust store. It + has no effect on Linux/macOS, where bash is already used. + +If ``weightslab se --force-ubuntu`` hangs with no output, WSL itself is not +responding (``wsl -e echo ok`` hangs too). Run ``wsl --shutdown`` and retry, +or drop ``--force-ubuntu`` to use PowerShell. + weightslab start ~~~~~~~~~~~~~~~~ @@ -58,22 +84,43 @@ weightslab start [--backend-host HOST] [--backend-port PORT] [--no-browser] [--certs] -Runs the UI natively from Python. +Runs the UI natively from Python: one process serves the bundled Weights +Studio page and proxies gRPC-Web to the training backend. Unsecured HTTP by +default. -``DIR`` *(positional, optional)* — establishes the experiment directory (its -checkpoints, logs, and ``notebook.ipynb`` live there). UI-only; it does not -start training on its own. +**Arguments** + +- ``DIR`` *(positional, optional)* — establishes the experiment directory (its + checkpoints, logs, and ``notebook.ipynb`` live there); created if missing. + Omit it to create a fresh ``./wl--`` directory. UI-only; it + does not start training on its own. +- ``--port`` *(int)* — UI HTTP port; see the resolution order below. +- ``--config`` *(file)* — experiment config (YAML) to read the UI port from. +- ``--host`` *(str)* — interface the UI binds to. Default: + ``$WEIGHTSLAB_UI_HOST``, else **0.0.0.0**. +- ``--backend-host`` *(str)* — backend gRPC host to proxy to. Default: + ``$GRPC_BACKEND_HOST``, else **localhost**. +- ``--backend-port`` *(int)* — backend gRPC port to proxy to. Default: + ``$GRPC_BACKEND_PORT``, else **50051**. +- ``--no-browser`` — don't open a browser tab. +- ``--certs`` — serve HTTPS, and use mTLS to the backend, with the certificates + in ``$WEIGHTSLAB_CERTS_DIR``, else ``~/.weightslab-certs`` (run + ``weightslab se`` first). ``~/.weightslab-certs`` is also used when the + variable points at a directory without certs. If no valid certificates are + found it logs a warning and serves plain HTTP. Port resolution order: 1. --port -2. ui_port from --config / WEIGHTSLAB_EXPERIMENT_CONFIG config file +2. ui_port from the config file: --config, else WEIGHTSLAB_EXPERIMENT_CONFIG, + else ./config.yaml, ./config.yml, ./experiment_config.yaml or + ./experiment_config.yml in the current directory 3. WL_LAST_UI_PORT 4. WEIGHTSLAB_UI_PORT (compatibility) 5. 8080 -If the chosen port is already in use, weightslab start falls back to a random -available port and logs it. +If the chosen port is already in use, or is the backend's gRPC port, +weightslab start picks a free port instead and logs it. Examples: @@ -117,7 +164,7 @@ the documented form. * - ``--clus`` - Clustering * - ``--gen`` - - Generation + - Image generation (reconstruction + contrastive, anomaly detection) * - ``--3d_det`` - 3D LiDAR point-cloud detection * - ``--2d_det`` @@ -149,9 +196,9 @@ One-level-at-a-time MNIST demos (four-way SDK approach — see weightslab start example --3d_det # 3D LiDAR detection weightslab example start --det # tolerant alias, same as `start example --det` -Then, in another terminal: ``weightslab start`` and open -``http://localhost:5173``. See :doc:`examples/index` for what each example -demonstrates. +Then, in another terminal, run ``weightslab start`` and open the URL it +prints (``http://localhost:8080`` by default). See :doc:`examples/index` for +what each example demonstrates. weightslab cli ~~~~~~~~~~~~~~ @@ -160,7 +207,15 @@ weightslab cli weightslab cli [--port PORT] [--host HOST] -Connects to a running experiment CLI server. +Opens an interactive console attached to a running experiment's CLI server. +The experiment must serve it (``wl.serve(serving_cli=True)``). + +- ``--port`` *(int)* — CLI server port. Default: auto-discover the running + experiment (it advertises its port on startup), else ``$CLI_PORT``. +- ``--host`` *(str)* — CLI server host. Default: the host the experiment + advertised, else ``$CLI_HOST``, else **localhost**. + +The console commands are listed under `Interactive CLI console`_ below. weightslab agent ~~~~~~~~~~~~~~~~ @@ -184,14 +239,15 @@ weightslab tunnel weightslab tunnel [ENDPOINT] [--listen-port N] [--listen-host H] [--remote-port N] Forwards a **remote** gRPC training backend to a **local** TCP port so the -Weights Studio UI — whose Envoy proxy dials ``localhost:50051`` — connects to -it as if it were local. This is what lets you **train on a remote machine (e.g. -Google Colab) and watch it live in Studio running on your laptop**: Colab has no -Docker daemon, so you run the UI locally and bridge the remote backend to it. +Weights Studio UI — whose ``weightslab start`` proxy dials ``localhost:50051`` +by default — connects to it as if it were local. This is what lets you **train +on a remote machine (e.g. Google Colab) and watch it live in Studio running on +your laptop**: you run the UI locally and bridge the remote backend to it. It is a raw byte forwarder (no protocol parsing) because the browser speaks -gRPC-Web to Envoy and Envoy speaks native HTTP/2 gRPC to its upstream — those -HTTP/2 frames must pass through untouched. Two consequences: +gRPC-Web to the ``weightslab start`` server, which speaks native HTTP/2 gRPC to +its upstream — those HTTP/2 frames must pass through untouched. Two +consequences: - The remote tunnel must be **raw TCP**, *not* an HTTP/gRPC-Web tunnel. A zero-signup option is `bore `_ with its free @@ -207,12 +263,12 @@ HTTP/2 frames must pass through untouched. Two consequences: stripped. Default: the ``WEIGHTSLAB_TUNNEL_ENDPOINT`` environment variable, so a bare ``weightslab tunnel`` works once that is exported. - ``--listen-port``, ``-p`` *(int)* — local port to expose. Default: **50051** - (the port the bundled Envoy upstream dials — leave it unless you changed - ``GRPC_BACKEND_PORT``). + (the port ``weightslab start`` proxies to by default — leave it unless you + pass ``--backend-port`` or set ``GRPC_BACKEND_PORT``). - ``--listen-host`` *(str)* — interface to bind. Default: **auto** — - ``127.0.0.1`` on Windows/macOS (Docker Desktop reaches host loopback via - ``host.docker.internal``), ``0.0.0.0`` on Linux (compose ``host-gateway`` - resolves to the bridge IP, which cannot reach a loopback-only listener). + ``127.0.0.1`` on Windows/macOS, ``0.0.0.0`` (all interfaces) on Linux. With + the UI on the same machine, ``--listen-host 127.0.0.1`` works on Linux too + and keeps the tunnel private. - ``--remote-port`` *(int)* — the remote port, when ``ENDPOINT`` has only a host and no ``:port``. @@ -237,7 +293,8 @@ HTTP/2 frames must pass through untouched. Two consequences: weightslab start # plaintext HTTP (default) weightslab tunnel bore.pub:12345 # the host:port bore printed - # 3) Open http://localhost:5173 — Studio streams live from Colab. + # 3) Open the URL `weightslab start` printed (http://localhost:8080 by + # default) — Studio streams live from Colab. .. note:: diff --git a/docs/weights_studio/configuration.rst b/docs/weights_studio/configuration.rst index bf16d7d1..3e9b7dd7 100644 --- a/docs/weights_studio/configuration.rst +++ b/docs/weights_studio/configuration.rst @@ -17,7 +17,8 @@ Backend environment variables (set before starting ``wl.serve()``) +----------------------------------+-------------------------+----------------------------------------------------+ | ``GRPC_TLS_REQUIRE_CLIENT_AUTH`` | ``0`` | ``1`` = require client mTLS certificate | +----------------------------------+-------------------------+----------------------------------------------------+ -| ``WEIGHTSLAB_CERTS_DIR`` | ``~/.weightslab-certs`` | Directory containing cert/key files | +| ``WEIGHTSLAB_CERTS_DIR`` | ``~/.weightslab-certs`` | Directory containing cert/key files; when it | +| | | holds none, ``~/.weightslab-certs`` is used | +----------------------------------+-------------------------+----------------------------------------------------+ | ``GRPC_AUTH_TOKEN`` | *(unset)* | Optional metadata-token auth (on top of mTLS) | +----------------------------------+-------------------------+----------------------------------------------------+ @@ -40,7 +41,8 @@ UI server environment variables (set before ``weightslab start``) +---------------------------+-------------------------+--------------------------------------------------+ | ``GRPC_BACKEND_PORT`` | ``50051`` | Backend gRPC port to proxy to | +---------------------------+-------------------------+--------------------------------------------------+ -| ``WEIGHTSLAB_CERTS_DIR`` | ``~/.weightslab-certs`` | Certs dir (read when ``--certs``) | +| ``WEIGHTSLAB_CERTS_DIR`` | ``~/.weightslab-certs`` | Certs dir (read when ``--certs``); when it | +| | | holds none, ``~/.weightslab-certs`` is used | +---------------------------+-------------------------+--------------------------------------------------+ | ``WEIGHTSLAB_OPENCODE_PORT`` | ``4096`` | Port the agent (OpenCode) server is started on; | | | | falls back to a free port if taken | diff --git a/docs/weights_studio/security.rst b/docs/weights_studio/security.rst index e812da61..0e321287 100644 --- a/docs/weights_studio/security.rst +++ b/docs/weights_studio/security.rst @@ -7,15 +7,26 @@ The default is plain HTTP (no cert files required, easiest for local dev). Do th weightslab se - Certificates are placed in ``~/.weightslab-certs`` - (or ``$WEIGHTSLAB_CERTS_DIR``). + Certificates are placed in ``$WEIGHTSLAB_CERTS_DIR``, else + ``~/.weightslab-certs``. Follow the printed instructions to export ``WEIGHTSLAB_CERTS_DIR`` globally. + On Windows, ``weightslab se`` runs the PowerShell script with the Windows + ``openssl`` and adds the dev CA to your user's trusted root certificates + (Windows asks you to confirm). To generate the certificates through WSL + (Ubuntu) instead, run:: + + weightslab se --force-ubuntu + + The WSL path does not install the CA into the Windows trust store. + 2. Start the UI in secure mode:: weightslab start --certs - ``--certs`` reads ``$WEIGHTSLAB_CERTS_DIR`` (single source of truth) and: + ``--certs`` reads ``$WEIGHTSLAB_CERTS_DIR`` (single source of truth). When + the variable is unset, or its directory has no certs, + ``~/.weightslab-certs`` is used instead. It then: - Serves HTTPS using ``ui-server.crt`` / ``ui-server.key`` - Presents ``ui-client.crt`` / ``ui-client.key`` to the backend (mTLS) diff --git a/docs/weights_studio_ui/more/configuration.rst b/docs/weights_studio_ui/more/configuration.rst index bf16d7d1..3e9b7dd7 100644 --- a/docs/weights_studio_ui/more/configuration.rst +++ b/docs/weights_studio_ui/more/configuration.rst @@ -17,7 +17,8 @@ Backend environment variables (set before starting ``wl.serve()``) +----------------------------------+-------------------------+----------------------------------------------------+ | ``GRPC_TLS_REQUIRE_CLIENT_AUTH`` | ``0`` | ``1`` = require client mTLS certificate | +----------------------------------+-------------------------+----------------------------------------------------+ -| ``WEIGHTSLAB_CERTS_DIR`` | ``~/.weightslab-certs`` | Directory containing cert/key files | +| ``WEIGHTSLAB_CERTS_DIR`` | ``~/.weightslab-certs`` | Directory containing cert/key files; when it | +| | | holds none, ``~/.weightslab-certs`` is used | +----------------------------------+-------------------------+----------------------------------------------------+ | ``GRPC_AUTH_TOKEN`` | *(unset)* | Optional metadata-token auth (on top of mTLS) | +----------------------------------+-------------------------+----------------------------------------------------+ @@ -40,7 +41,8 @@ UI server environment variables (set before ``weightslab start``) +---------------------------+-------------------------+--------------------------------------------------+ | ``GRPC_BACKEND_PORT`` | ``50051`` | Backend gRPC port to proxy to | +---------------------------+-------------------------+--------------------------------------------------+ -| ``WEIGHTSLAB_CERTS_DIR`` | ``~/.weightslab-certs`` | Certs dir (read when ``--certs``) | +| ``WEIGHTSLAB_CERTS_DIR`` | ``~/.weightslab-certs`` | Certs dir (read when ``--certs``); when it | +| | | holds none, ``~/.weightslab-certs`` is used | +---------------------------+-------------------------+--------------------------------------------------+ | ``WEIGHTSLAB_OPENCODE_PORT`` | ``4096`` | Port the agent (OpenCode) server is started on; | | | | falls back to a free port if taken | diff --git a/docs/weights_studio_ui/more/security.rst b/docs/weights_studio_ui/more/security.rst index e812da61..0e321287 100644 --- a/docs/weights_studio_ui/more/security.rst +++ b/docs/weights_studio_ui/more/security.rst @@ -7,15 +7,26 @@ The default is plain HTTP (no cert files required, easiest for local dev). Do th weightslab se - Certificates are placed in ``~/.weightslab-certs`` - (or ``$WEIGHTSLAB_CERTS_DIR``). + Certificates are placed in ``$WEIGHTSLAB_CERTS_DIR``, else + ``~/.weightslab-certs``. Follow the printed instructions to export ``WEIGHTSLAB_CERTS_DIR`` globally. + On Windows, ``weightslab se`` runs the PowerShell script with the Windows + ``openssl`` and adds the dev CA to your user's trusted root certificates + (Windows asks you to confirm). To generate the certificates through WSL + (Ubuntu) instead, run:: + + weightslab se --force-ubuntu + + The WSL path does not install the CA into the Windows trust store. + 2. Start the UI in secure mode:: weightslab start --certs - ``--certs`` reads ``$WEIGHTSLAB_CERTS_DIR`` (single source of truth) and: + ``--certs`` reads ``$WEIGHTSLAB_CERTS_DIR`` (single source of truth). When + the variable is unset, or its directory has no certs, + ``~/.weightslab-certs`` is used instead. It then: - Serves HTTPS using ``ui-server.crt`` / ``ui-server.key`` - Presents ``ui-client.crt`` / ``ui-client.key`` to the backend (mTLS) diff --git a/tests/backend/test_cli.py b/tests/backend/test_cli.py index 5a164915..ae1a4758 100644 --- a/tests/backend/test_cli.py +++ b/tests/backend/test_cli.py @@ -17,6 +17,7 @@ from unittest.mock import patch, MagicMock from weightslab.cli import ( + _EXAMPLES, _convert_to_git_bash_path, _ensure_scripts_executable, _generate_certs_with_fallback, @@ -25,6 +26,7 @@ _install_example_requirements, _make_executable, _resolve_experiment_dir, + _resolve_ui_port, example_start, main, ui_secure_environment, @@ -87,9 +89,20 @@ def test_unix_path_passthrough(self): ) +class TestPersistCertsDir(unittest.TestCase): + @patch("weightslab.cli.subprocess.run") + def test_refuses_relative_path(self, mock_run): + from weightslab.cli import _persist_certs_dir + with patch("weightslab.cli.Path.home") as mock_home: + _persist_certs_dir("") + mock_run.assert_not_called() # no setx + mock_home.assert_not_called() # no ~/.bashrc write + + class TestCertGeneration(unittest.TestCase): + @patch("weightslab.cli._is_windows", return_value=False) @patch("weightslab.cli._run_shell_script", return_value=0) - def test_generate_certs_forwards_certs_dir(self, mock_shell): + def test_generate_certs_forwards_certs_dir(self, mock_shell, _mock_win): rc = _generate_certs_with_fallback(force_certs=False, certs_dir="/custom/certs") self.assertEqual(rc, 0) env_vars = mock_shell.call_args.args[2] @@ -97,13 +110,79 @@ def test_generate_certs_forwards_certs_dir(self, mock_shell): self.assertIn("WEIGHTSLAB_CERTS_DIR", env_vars) self.assertIn("custom/certs", env_vars["WEIGHTSLAB_CERTS_DIR"].replace("\\", "/")) + @patch("weightslab.cli._is_windows", return_value=False) @patch("weightslab.cli._run_shell_script", return_value=0) - def test_generate_certs_without_dir_passes_no_env(self, mock_shell): + def test_generate_certs_without_dir_passes_no_env(self, mock_shell, _mock_win): _generate_certs_with_fallback(force_certs=False) self.assertIsNone(mock_shell.call_args.args[2]) + @patch("weightslab.cli._is_windows", return_value=True) + @patch("weightslab.cli._run_shell_script", return_value=0) + @patch("weightslab.cli._run_powershell_script", return_value=0) + def test_windows_defaults_to_powershell_with_native_path(self, mock_ps, mock_shell, _mock_win): + certs_dir = r"C:\Users\testuser\.weightslab-certs" + rc = _generate_certs_with_fallback(force_certs=True, certs_dir=certs_dir) + self.assertEqual(rc, 0) + mock_shell.assert_not_called() + script, args, env_vars = mock_ps.call_args.args + self.assertTrue(script.endswith("generate-certs-auth-token.ps1")) + self.assertEqual(args, ["-ForceCreateCerts"]) + # PowerShell gets the host path, not the WSL /mnt/c form. + self.assertEqual(env_vars, {"WEIGHTSLAB_CERTS_DIR": certs_dir}) + + @patch("weightslab.cli._is_windows", return_value=True) + @patch("weightslab.cli._run_shell_script", return_value=0) + @patch("weightslab.cli._run_powershell_script", return_value=1) + def test_windows_falls_back_to_bash_when_powershell_fails(self, mock_ps, mock_shell, _mock_win): + rc = _generate_certs_with_fallback(certs_dir=r"C:\Users\testuser\.weightslab-certs") + self.assertEqual(rc, 0) + mock_ps.assert_called_once() + mock_shell.assert_called_once() + self.assertEqual(mock_shell.call_args.args[2], + {"WEIGHTSLAB_CERTS_DIR": "/mnt/c/Users/testuser/.weightslab-certs"}) + + @patch("weightslab.cli._is_windows", return_value=True) + @patch("weightslab.cli._run_shell_script", return_value=0) + @patch("weightslab.cli._run_powershell_script", return_value=0) + def test_windows_force_ubuntu_uses_bash_only(self, mock_ps, mock_shell, _mock_win): + rc = _generate_certs_with_fallback( + force_certs=True, certs_dir=r"C:\Users\testuser\.weightslab-certs", force_ubuntu=True) + self.assertEqual(rc, 0) + mock_ps.assert_not_called() + script, args, env_vars = mock_shell.call_args.args + self.assertTrue(script.endswith("generate-certs-auth-token.sh")) + self.assertEqual(args, ["--force-create-certs"]) + self.assertEqual(env_vars, {"WEIGHTSLAB_CERTS_DIR": "/mnt/c/Users/testuser/.weightslab-certs"}) + + @patch("weightslab.cli._is_windows", return_value=True) + @patch("weightslab.cli._run_shell_script", return_value=2) + @patch("weightslab.cli._run_powershell_script", return_value=0) + def test_windows_force_ubuntu_does_not_fall_back_to_powershell(self, mock_ps, _mock_shell, _mock_win): + rc = _generate_certs_with_fallback(force_ubuntu=True) + self.assertEqual(rc, 2) + mock_ps.assert_not_called() + + @patch("weightslab.cli._is_windows", return_value=False) + @patch("weightslab.cli._run_shell_script", return_value=0) + @patch("weightslab.cli._run_powershell_script", return_value=0) + def test_force_ubuntu_is_noop_off_windows(self, mock_ps, mock_shell, _mock_win): + self.assertEqual(_generate_certs_with_fallback(force_ubuntu=True), 0) + self.assertEqual(_generate_certs_with_fallback(force_ubuntu=False), 0) + mock_ps.assert_not_called() + self.assertEqual(mock_shell.call_count, 2) + class TestUiSecureEnvironment(unittest.TestCase): + def setUp(self): + # `se` exports WEIGHTSLAB_CERTS_DIR into os.environ and reads it as its + # default dir: sandbox it so tests neither leak it into later tests nor + # write a real token into ~/.weightslab-certs. + tmp = tempfile.TemporaryDirectory() + self.addCleanup(tmp.cleanup) + env = patch.dict(os.environ, {"WEIGHTSLAB_CERTS_DIR": tmp.name}) + env.start() + self.addCleanup(env.stop) + @patch("weightslab.cli.CertAuthManager") @patch("weightslab.cli._generate_certs_with_fallback", return_value=0) def test_ui_secure_environment_success(self, mock_gen_certs, mock_cert_manager): @@ -114,7 +193,8 @@ def test_ui_secure_environment_success(self, mock_gen_certs, mock_cert_manager): with self.assertLogs("weightslab.cli", level="INFO") as log_context: ui_secure_environment(argparse.Namespace(force_certs=False)) self.assertTrue(any("Certificates generated successfully" in m for m in log_context.output)) - mock_gen_certs.assert_called_once_with(force_certs=False, certs_dir=mgr.certs_dir) + mock_gen_certs.assert_called_once_with( + force_certs=False, certs_dir=mgr.certs_dir, force_ubuntu=False) mgr.certs_dir.mkdir.assert_called_once() self.assertTrue(any("WEIGHTSLAB_CERTS_DIR exported" in m for m in log_context.output)) @@ -124,6 +204,46 @@ def test_ui_secure_environment_force_certs(self, mock_gen_certs): ui_secure_environment(argparse.Namespace(force_certs=True)) self.assertTrue(mock_gen_certs.call_args.kwargs["force_certs"]) + @patch("weightslab.cli._generate_certs_with_fallback", return_value=0) + def test_ui_secure_environment_force_ubuntu(self, mock_gen_certs): + with patch.dict(os.environ, {}, clear=False): + ui_secure_environment(argparse.Namespace(force_certs=False, force_ubuntu=True)) + self.assertTrue(mock_gen_certs.call_args.kwargs["force_ubuntu"]) + + @patch("weightslab.cli.CertAuthManager") + @patch("weightslab.cli._generate_certs_with_fallback", return_value=0) + def test_ui_secure_environment_uses_env_certs_dir(self, _mock_gen_certs, mock_cert_manager): + custom = str(Path(tempfile.gettempdir()) / "custom-certs") + with patch.dict(os.environ, {"WEIGHTSLAB_CERTS_DIR": custom}): + ui_secure_environment(argparse.Namespace(force_certs=False)) + self.assertEqual(mock_cert_manager.call_args.kwargs["certs_dir"], custom) + + @patch("weightslab.cli.CertAuthManager") + @patch("weightslab.cli._generate_certs_with_fallback", return_value=0) + def test_ui_secure_environment_ignores_relative_env(self, _mock_gen_certs, mock_cert_manager): + """A bogus value (e.g. a leaked mock repr) falls back to the default dir.""" + bogus = "" + with patch.dict(os.environ, {"WEIGHTSLAB_CERTS_DIR": bogus}): + ui_secure_environment(argparse.Namespace(force_certs=False)) + self.assertIsNone(mock_cert_manager.call_args.kwargs["certs_dir"]) + + @patch("weightslab.cli.CertAuthManager") + @patch("weightslab.cli._generate_certs_with_fallback", return_value=0) + def test_ui_secure_environment_positional_beats_env(self, _mock_gen_certs, mock_cert_manager): + with tempfile.TemporaryDirectory() as tmp: + with patch.dict(os.environ, {"WEIGHTSLAB_CERTS_DIR": "/custom/certs"}): + ui_secure_environment(argparse.Namespace(force_certs=False, certs_dir=tmp)) + self.assertEqual(mock_cert_manager.call_args.kwargs["certs_dir"], + str(Path(tmp).resolve())) + + @patch("weightslab.cli.CertAuthManager") + @patch("weightslab.cli._generate_certs_with_fallback", return_value=0) + def test_ui_secure_environment_default_dir_without_env(self, _mock_gen_certs, mock_cert_manager): + env = {k: v for k, v in os.environ.items() if k != "WEIGHTSLAB_CERTS_DIR"} + with patch.dict(os.environ, env, clear=True): + ui_secure_environment(argparse.Namespace(force_certs=False)) + self.assertIsNone(mock_cert_manager.call_args.kwargs["certs_dir"]) + @patch("weightslab.cli._generate_certs_with_fallback", return_value=1) def test_ui_secure_environment_cert_failure(self, _mock_gen_certs): with self.assertRaises(SystemExit) as ctx: @@ -204,6 +324,26 @@ def test_start_survives_telemetry_failure(self, mock_serve, _mock_ping): mock_serve.assert_called_once() +class TestResolveUiPort(unittest.TestCase): + @patch("weightslab.cli._load_ui_port_from_experiment_config", return_value=None) + def test_default_is_8080(self, _mock_cfg): + env = {k: v for k, v in os.environ.items() + if k not in ("WL_LAST_UI_PORT", "WEIGHTSLAB_UI_PORT")} + with patch.dict(os.environ, env, clear=True): + self.assertEqual(_resolve_ui_port(argparse.Namespace(port=None, config=None)), + (8080, "default")) + + @patch("weightslab.cli._load_ui_port_from_experiment_config", return_value=None) + def test_env_order(self, _mock_cfg): + ns = argparse.Namespace(port=None, config=None) + with patch.dict(os.environ, {"WL_LAST_UI_PORT": "9001", "WEIGHTSLAB_UI_PORT": "9002"}): + self.assertEqual(_resolve_ui_port(ns), (9001, "WL_LAST_UI_PORT")) + env = {k: v for k, v in os.environ.items() if k != "WL_LAST_UI_PORT"} + env["WEIGHTSLAB_UI_PORT"] = "9002" + with patch.dict(os.environ, env, clear=True): + self.assertEqual(_resolve_ui_port(ns), (9002, "WEIGHTSLAB_UI_PORT")) + + class TestExperimentDir(unittest.TestCase): """`weightslab start [DIR]` establishes the experiment directory (root_log_dir).""" @@ -292,6 +432,12 @@ def test_example_start_errors_when_missing(self, _mock_dir): def test_example_dir_points_at_bundled_example(self): self.assertTrue((_get_example_dir("wl-classification") / "main.py").exists()) + def test_every_example_flag_points_at_a_bundled_main_py(self): + for kind, (dir_name, _label, category) in _EXAMPLES.items(): + with self.subTest(kind=kind): + self.assertTrue((_get_example_dir(dir_name, category) / "main.py").exists(), + f"--{kind} -> examples/{category}/{dir_name}/main.py is missing") + @patch("weightslab.cli.subprocess.run") def test_example_start_seg_runs_segmentation(self, mock_run): mock_run.return_value = MagicMock(returncode=0) @@ -329,6 +475,18 @@ def test_main_dispatches_se(self, mock_se): main() mock_se.assert_called_once() + @patch("weightslab.cli.ui_secure_environment") + def test_main_se_force_ubuntu_flag(self, mock_se): + with patch("sys.argv", ["weightslab", "se", "--force-ubuntu"]): + main() + self.assertTrue(mock_se.call_args.args[0].force_ubuntu) + + @patch("weightslab.cli.ui_secure_environment") + def test_main_se_defaults_force_ubuntu_off(self, mock_se): + with patch("sys.argv", ["weightslab", "se"]): + main() + self.assertFalse(mock_se.call_args.args[0].force_ubuntu) + @patch("weightslab.cli.ui_start_native") def test_main_bare_start_launches_native_ui(self, mock_start): with patch("sys.argv", ["weightslab", "start"]): @@ -374,6 +532,21 @@ def test_help_shows_new_command_set(self): self.assertIn("start", out) self.assertIn("start example", out) + def test_help_lists_every_command_and_hides_alias(self): + out = self._capture_main(["weightslab", "help"]) + self.assertNotIn("==SUPPRESS==", out) + for cmd in ("se [CERTS_DIR]", "start [DIR]", "start example", "cli", "tunnel", + "export", "agent init"): + with self.subTest(cmd=cmd): + self.assertIn(cmd, out) + + def test_every_subcommand_has_help(self): + for argv in (["se"], ["start"], ["start", "example"], ["example", "start"], + ["cli"], ["tunnel"], ["export"], ["agent"], ["agent", "init"], ["help"]): + with self.subTest(argv=argv): + out = self._capture_main(["weightslab", *argv, "--help"]) + self.assertIn("usage: weightslab", out) + if __name__ == "__main__": unittest.main() diff --git a/tests/general/test_cli.py b/tests/general/test_cli.py index a63801dc..64dd619b 100644 --- a/tests/general/test_cli.py +++ b/tests/general/test_cli.py @@ -563,10 +563,10 @@ def test_last_ui_port_env_fallback(self): self.assertEqual(port, 61236) self.assertEqual(source, "WL_LAST_UI_PORT") - def test_default_is_50051(self): + def test_default_is_8080(self): args = argparse.Namespace(port=None, config=None) port, source = wl_cli._resolve_ui_port(args) - self.assertEqual(port, 50051) + self.assertEqual(port, 8080) self.assertEqual(source, "default") diff --git a/tests/test_secure_communication.py b/tests/test_secure_communication.py index 75c38f19..5a69ef1f 100644 --- a/tests/test_secure_communication.py +++ b/tests/test_secure_communication.py @@ -75,13 +75,12 @@ def test_auth_environment_variables_disabled(self): assert "GRPC_AUTH_TOKEN" not in env_vars def test_from_env_or_default(self): - with tempfile.TemporaryDirectory() as tmpdir: - os.environ["WEIGHTSLAB_CERTS_DIR"] = tmpdir - try: + # Isolated from the real ~/.weightslab-certs, which may hold certs. + with tempfile.TemporaryDirectory() as tmpdir, tempfile.TemporaryDirectory() as home: + with patch("weightslab.security.cert_auth_manager._get_user_profile", return_value=home), \ + patch.dict(os.environ, {"WEIGHTSLAB_CERTS_DIR": tmpdir}): manager = CertAuthManager.from_env_or_default() assert manager.certs_dir == Path(tmpdir) - finally: - del os.environ["WEIGHTSLAB_CERTS_DIR"] @patch("weightslab.security.cert_auth_manager.subprocess.run") def test_generate_certs_success(self, mock_run): @@ -130,3 +129,62 @@ def test_secure_init_skip_env_var(self): if __name__ == "__main__": pytest.main([__file__, "-v"]) + + +def _write_cert_set(directory): + Path(directory).mkdir(parents=True, exist_ok=True) + for name in ("backend-server.crt", "backend-server.key", "ca.crt"): + (Path(directory) / name).write_text("x") + + +class TestCertsDirResolution: + """$WEIGHTSLAB_CERTS_DIR first, then ~/.weightslab-certs.""" + + @pytest.fixture(autouse=True) + def _home(self, tmp_path): + from weightslab.security import cert_auth_manager as cam + cam._WARNED_CERTS_DIRS.clear() + home = tmp_path / "home" + home.mkdir() + self.default_dir = home / ".weightslab-certs" + self.env_dir = tmp_path / "custom-certs" + with patch.object(cam, "_get_user_profile", return_value=str(home)): + env = {k: v for k, v in os.environ.items() if k != "WEIGHTSLAB_CERTS_DIR"} + with patch.dict(os.environ, env, clear=True): + yield + + def test_unset_uses_default(self): + assert CertAuthManager.from_env_or_default().certs_dir == self.default_dir + + def test_empty_uses_default(self): + with patch.dict(os.environ, {"WEIGHTSLAB_CERTS_DIR": " "}): + assert CertAuthManager.from_env_or_default().certs_dir == self.default_dir + + def test_env_dir_with_certs_wins(self): + _write_cert_set(self.env_dir) + _write_cert_set(self.default_dir) + with patch.dict(os.environ, {"WEIGHTSLAB_CERTS_DIR": str(self.env_dir)}): + assert CertAuthManager.from_env_or_default().certs_dir == self.env_dir + + def test_falls_back_to_default_when_env_dir_has_no_certs(self, caplog): + _write_cert_set(self.default_dir) + with patch.dict(os.environ, {"WEIGHTSLAB_CERTS_DIR": str(self.env_dir)}): + assert CertAuthManager.from_env_or_default().certs_dir == self.default_dir + assert "No certs in WEIGHTSLAB_CERTS_DIR" in caplog.text + + def test_keeps_env_dir_when_neither_has_certs(self): + with patch.dict(os.environ, {"WEIGHTSLAB_CERTS_DIR": str(self.env_dir)}): + assert CertAuthManager.from_env_or_default().certs_dir == self.env_dir + + def test_relative_value_is_ignored(self, caplog): + _write_cert_set(self.default_dir) + bogus = "" + with patch.dict(os.environ, {"WEIGHTSLAB_CERTS_DIR": bogus}): + assert CertAuthManager.from_env_or_default().certs_dir == self.default_dir + assert "not an absolute path" in caplog.text + + def test_warning_logged_once(self, caplog): + with patch.dict(os.environ, {"WEIGHTSLAB_CERTS_DIR": "relative/dir"}): + CertAuthManager.from_env_or_default() + CertAuthManager.from_env_or_default() + assert caplog.text.count("not an absolute path") == 1 diff --git a/tests/trainer/services/test_trainer_services_server.py b/tests/trainer/services/test_trainer_services_server.py index 9b8cc41d..af97659c 100644 --- a/tests/trainer/services/test_trainer_services_server.py +++ b/tests/trainer/services/test_trainer_services_server.py @@ -410,9 +410,13 @@ def start(self): return fake_server def _with_grpc_env(self, host, port, fn): - saved = {k: os.environ.get(k) for k in ("GRPC_BACKEND_HOST", "GRPC_BACKEND_PORT")} + # Pin plaintext: `import weightslab` sets GRPC_TLS_ENABLED=1 when the + # machine has certs in $WEIGHTSLAB_CERTS_DIR / ~/.weightslab-certs. + saved = {k: os.environ.get(k) + for k in ("GRPC_BACKEND_HOST", "GRPC_BACKEND_PORT", "GRPC_TLS_ENABLED")} os.environ["GRPC_BACKEND_HOST"] = host os.environ["GRPC_BACKEND_PORT"] = port + os.environ["GRPC_TLS_ENABLED"] = "0" try: return fn() finally: diff --git a/weightslab/AGENTS.md b/weightslab/AGENTS.md index f31ebd7b..42851eb5 100644 --- a/weightslab/AGENTS.md +++ b/weightslab/AGENTS.md @@ -77,7 +77,8 @@ whole API surface (reactive signals, group signals, the Ultralytics mixin, etc. aren't in the `.rst` docs; the examples are the primary source). TLS/UI deploy details: `weightslab/docs/weights_studio.rst`. TLS is opt-in: -`weightslab se` once, then `weightslab start --certs`. +`weightslab se` once, then `weightslab start --certs`. Windows: `se` uses the +PowerShell script + Windows `openssl`; `--force-ubuntu` uses WSL bash instead. --- @@ -193,7 +194,7 @@ from a watched object: - `wl.save_group_signals(signals={...}, group_ids=[...], origin="train_loader")` — one row per group, for pairwise values (e.g. contrastive loss) that can't map to a single sample. Needs a dataset that emits a `group_id` in its metadata - (`PyTorch/wl-generation`). + (`PyTorch/wl-image-generation`). - `wl.trajectory_stats(values)` / `wl.classify_loss_shape(values)` — building blocks behind the loss-shape tag (§3.6); call directly only for a custom classifier. @@ -285,7 +286,7 @@ The automatic `tag:loss_shape` tag (§3.6) uses these same primitives. | LiDAR detection (2D/3D) | `Usecases/wl-{2d,3d}-lidar-detection` | `task_type="detection_pointcloud"`; 3D adds `render_thumbnail_2d`. | | Tabular / feature vectors | `PyTorch/wl-fraud-detection` | No `task_type`; `preload_labels=True, preload_metadata=True`; see §3.10 for a headless verification script. | | Embedding / clustering | `PyTorch/wl-clustering` (+ `face/model.py`) | `watch_or_edit` calls live inside the model wrapper, not `main.py`; open-ended loop. | -| Paired/contrastive samples, group-level signals | `PyTorch/wl-generation` | `wl.save_group_signals`; dataset emits 2 rows per item via a `uids` metadata key. | +| Paired/contrastive samples, group-level signals | `PyTorch/wl-image-generation` | `wl.save_group_signals`; dataset emits 2 rows per item via a `uids` metadata key. | | Reactive signals / custom loss-shape tagging | `Usecases/wl-classification-signals_shape_classification`, `Usecases/ws-signals-mnist` | §3.6; the latter is the minimal variant with no custom classifier. | | PyTorch Lightning | `Lightning/wl-classification` | Same `watch_or_edit` calls as plain PyTorch; guards wrap `training_step`/`validation_step` bodies; `Trainer(log_every_n_steps=0, enable_checkpointing=False, logger=False)`. | | Ultralytics YOLO (detect/segment) | `Ultralytics/wl-detection` | Don't call `watch_or_edit` for model/optimizer/data/loss/metric — pass `trainer=wl.WLAwareTrainer` (or `wl.WLAwareSegmentationTrainer`) to `YOLO(...).train(...)`. It wires everything via UL callbacks; you only watch the run config as `flag="hyperparameters"`. | @@ -313,7 +314,7 @@ Authoritative reference: `weightslab/docs/configuration.rst`. High-signal ones: | `GRPC_BACKEND_HOST`/`PORT` | `0.0.0.0`/`50051` | Backend gRPC bind address. | | `GRPC_TLS_ENABLED` | `0` | TLS on the gRPC socket; set with `weightslab start --certs`. | | `GRPC_TLS_REQUIRE_CLIENT_AUTH` | `0` | mTLS; must match `--certs`. | -| `WEIGHTSLAB_CERTS_DIR` | `~/.weightslab-certs` | Cert lookup — single source of truth. | +| `WEIGHTSLAB_CERTS_DIR` | `~/.weightslab-certs` | Cert lookup — single source of truth; falls back to `~/.weightslab-certs` when unset/relative/without certs. | | `GRPC_AUTH_TOKEN` | unset | Optional token auth on top of mTLS. | | `GRPC_MAX_MESSAGE_BYTES` | `268435456` | Raise if large tensors/images fail to transfer. | | `WEIGHTSLAB_DISABLE_WATCHDOGS` | `0` | Set `1` when breakpoint-debugging (§5). | @@ -345,6 +346,11 @@ backend serving on `0.0.0.0:50051`; (2) `weightslab start` running, browser reaches `:8080`; (3) TLS mismatch if using `--certs` — run `weightslab se` first, export `WEIGHTSLAB_CERTS_DIR` (or drop TLS: omit `--certs`, `GRPC_TLS_ENABLED=0`). +**`weightslab se` hangs with no output (Windows).** Only on the WSL path +(`--force-ubuntu`, or the fallback after PowerShell fails): a stuck WSL distro +blocks forever (output captured, no timeout). `wsl -e echo ok` hangs too → +`wsl --shutdown`, or drop `--force-ubuntu`. + **Env var change not taking effect.** `VITE_*` → rebuild frontend. `WS_*`/`BB_*`/`ENABLE_*` → restart `weightslab start` + reload tab. diff --git a/weightslab/cli.py b/weightslab/cli.py index 65a592e1..acd951d6 100644 --- a/weightslab/cli.py +++ b/weightslab/cli.py @@ -9,6 +9,8 @@ gRPC auth token) in $WEIGHTSLAB_CERTS_DIR. * ``weightslab cli`` — attach a terminal to a running experiment. * ``weightslab tunnel`` — forward a remote gRPC backend to a local port. + * ``weightslab export`` — export annotations to CVAT / Label Studio / V7. + * ``weightslab agent init`` — provision and sign in to the OpenCode assistant. """ import argparse @@ -22,7 +24,7 @@ import yaml -from weightslab.security import CertAuthManager +from weightslab.security import CertAuthManager, env_certs_dir from weightslab.tunnel import DEFAULT_LISTEN_PORT from weightslab.components.experiment_naming import generate_experiment_name as _generate_experiment_name @@ -120,7 +122,7 @@ def _resolve_ui_port(args) -> tuple[int, str]: if compat_env_port is not None: return compat_env_port, "WEIGHTSLAB_UI_PORT" - return 50051, "default" + return 8080, "default" def _persist_certs_dir(certs_dir_str: str) -> None: @@ -131,6 +133,9 @@ def _persist_certs_dir(certs_dir_str: str) -> None: Linux/macOS — appends an export line to ~/.bashrc (idempotent) and prints the source command for the current session. """ + if not Path(certs_dir_str).is_absolute(): + logger.warning(f"Not persisting WEIGHTSLAB_CERTS_DIR={certs_dir_str!r}: not an absolute path.") + return export_line = f'export WEIGHTSLAB_CERTS_DIR="{certs_dir_str}"' if _is_windows(): result = subprocess.run( @@ -176,12 +181,16 @@ def _banner() -> str: _EPILOG = """\ commands: - se Set up the secure environment: generate TLS - certificates + a gRPC auth token in - ~/.weightslab-certs. Then set WEIGHTSLAB_CERTS_DIR + se [CERTS_DIR] Set up the secure environment: generate TLS + certificates + a gRPC auth token in CERTS_DIR + (default: $WEIGHTSLAB_CERTS_DIR, else + ~/.weightslab-certs). Then set WEIGHTSLAB_CERTS_DIR (the single source of truth) so the backend + new shells find them. --force-certs regenerate even if certs exist + --force-ubuntu Windows only: generate with the + bash script through WSL/Ubuntu + instead of PowerShell (default) start [DIR] Start the Weights Studio UI natively — no Docker. Serves the bundled SPA and proxies gRPC-Web to a @@ -193,12 +202,14 @@ def _banner() -> str: fresh randomly-named dir under the current folder. (UI-only: run training separately with a main.py that points its root_log_dir at the same directory.) - --port PORT UI HTTP port (default 8080) + --port PORT UI HTTP port (default 8080) --config FILE experiment config file used to read ui_port + --host HOST UI bind host (default 0.0.0.0) --backend-port PORT backend gRPC port (default 50051) --backend-host HOST backend gRPC host (default localhost) --no-browser don't open a browser - --certs HTTPS + mTLS from $WEIGHTSLAB_CERTS_DIR + --certs HTTPS + mTLS from $WEIGHTSLAB_CERTS_DIR, + else ~/.weightslab-certs start example Run a bundled PyTorch example (foreground; stop with Ctrl+C). Installs the example's requirements first, @@ -207,7 +218,7 @@ def _banner() -> str: --seg segmentation example --det detection example --clus clustering example - --gen generation example + --gen image-generation example --3d_det 3D LiDAR point-cloud detection example --2d_det 2D LiDAR point-cloud detection example One-level-at-a-time MNIST demos (four-way SDK approach): @@ -247,9 +258,15 @@ def _banner() -> str: --host H backend host (default: 127.0.0.1) --port N backend gRPC port (default: 50051) + agent init Provision the OpenCode AI assistant (downloads the + binary, no Node.js needed), then run its interactive + sign-in so the Weights Studio agent is ready to use. + --provision-only download/verify only; skip sign-in + examples: weightslab se # one-time secure setup (then export WEIGHTSLAB_CERTS_DIR) weightslab se --force-certs # regenerate the certs + weightslab se --force-ubuntu # Windows: generate through WSL instead of PowerShell weightslab start # launch the UI (unsecured HTTP, default) at :8080 # (creates a fresh ./wl- experiment dir) weightslab start ./exp/mnist_opt/ # use (or create) this experiment directory @@ -267,6 +284,7 @@ def _banner() -> str: weightslab export --format cvat # export all annotations to CVAT XML in "." weightslab export -f v7 out/ --origin val_loader # V7/Darwin, val split only, into out/ weightslab export -f cvat --tag ToReview # only samples tagged ToReview, for relabeling + weightslab agent init # set up the AI assistant (OpenCode) once """ @@ -430,15 +448,13 @@ def _run_shell_script(script_path: str, args: list = None, env_vars: dict = None return 1 -def _generate_certs_with_fallback(force_certs: bool = False, certs_dir=None) -> int: - """Try shell script first, fall back to PowerShell on Windows if it fails. +def _generate_certs_bash(force_certs: bool, certs_dir) -> int: + """Run generate-certs-auth-token.sh (through WSL when on Windows).""" + cert_script = str(_get_cert_script()) + if not Path(cert_script).exists(): + logger.error(f"Shell script not found: {cert_script}") + return 1 - ``certs_dir`` is forwarded to the generation scripts as ``WEIGHTSLAB_CERTS_DIR`` - so certs land in the single source-of-truth directory (the scripts default to - ``~/.weightslab-certs`` when it is not provided). The shell script receives a - POSIX absolute path (``/mnt/c/...`` on Windows/WSL); PowerShell receives the - host-native path (``C:/...``). - """ env_vars = None if certs_dir is not None: # Shell scripts (bash/WSL) need a POSIX-style absolute path. @@ -448,37 +464,51 @@ def _generate_certs_with_fallback(force_certs: bool = False, certs_dir=None) -> bash_certs_dir = str(certs_dir) env_vars = {'WEIGHTSLAB_CERTS_DIR': bash_certs_dir} - cert_script = str(_get_cert_script()) - if not Path(cert_script).exists(): - logger.warning(f"Shell script not found: {cert_script}") - else: - script_args = [] - if force_certs: - script_args.append('--force-create-certs') + script_args = ['--force-create-certs'] if force_certs else [] + logger.info("Attempting certificate generation with shell script...") + return _run_shell_script(cert_script, script_args, env_vars) + + +def _generate_certs_powershell(force_certs: bool, certs_dir) -> int: + """Run generate-certs-auth-token.ps1 with the host openssl (Windows only).""" + cert_script_ps1 = str(_get_cert_script_ps1()) + if not Path(cert_script_ps1).exists(): + logger.error(f"PowerShell script not found: {cert_script_ps1}") + return 1 - logger.info("Attempting certificate generation with shell script...") - exit_code = _run_shell_script(cert_script, script_args, env_vars) + # Host-native path: PowerShell would resolve the /mnt/c/... form to C:\mnt\c\... + env_vars = {'WEIGHTSLAB_CERTS_DIR': str(certs_dir)} if certs_dir is not None else None + script_args = ['-ForceCreateCerts'] if force_certs else [] + logger.info("Attempting certificate generation with PowerShell script...") + return _run_powershell_script(cert_script_ps1, script_args, env_vars) + + +def _generate_certs_with_fallback(force_certs: bool = False, certs_dir=None, + force_ubuntu: bool = False) -> int: + """Generate the dev certs with the script native to this OS. + + On Windows the PowerShell script runs first (host openssl, no WSL) and the + bash script is only a fallback if it fails. ``force_ubuntu`` skips PowerShell + and runs only the bash script, i.e. through WSL/Ubuntu on Windows. Elsewhere + only the bash script applies, so ``force_ubuntu`` changes nothing. + + ``certs_dir`` is forwarded to the generation scripts as ``WEIGHTSLAB_CERTS_DIR`` + so certs land in the single source-of-truth directory (the scripts default to + ``~/.weightslab-certs`` when it is not provided). The shell script receives a + POSIX absolute path (``/mnt/c/...`` on Windows/WSL); PowerShell receives the + host-native path (``C:/...``). + """ + if _is_windows() and not force_ubuntu: + exit_code = _generate_certs_powershell(force_certs, certs_dir) if exit_code == 0: return 0 - logger.warning(f"Shell script failed (exit code {exit_code})") + logger.warning(f"PowerShell script failed (exit code {exit_code})") + logger.info("Falling back to bash (WSL) for certificate generation...") - # Fallback to PowerShell on Windows - if _is_windows(): - logger.info("Falling back to PowerShell for certificate generation...") - cert_script_ps1 = str(_get_cert_script_ps1()) - if not Path(cert_script_ps1).exists(): - logger.error(f"PowerShell script not found: {cert_script_ps1}") - return 1 - - script_args = [] - if force_certs: - script_args.append('-ForceCreateCerts') - - exit_code = _run_powershell_script(cert_script_ps1, script_args, env_vars) - return exit_code - else: - logger.error("Neither shell nor PowerShell script could generate certificates") - return 1 + exit_code = _generate_certs_bash(force_certs, certs_dir) + if exit_code != 0: + logger.warning(f"Shell script failed (exit code {exit_code})") + return exit_code def ui_secure_environment(args): @@ -495,17 +525,24 @@ def ui_secure_environment(args): _ensure_scripts_executable() force_certs = getattr(args, "force_certs", False) + force_ubuntu = getattr(args, "force_ubuntu", False) no_auth = getattr(args, "no_auth", False) certs_dir = getattr(args, "certs_dir", None) if certs_dir: # Absolute path so Windows Python, WSL bash and the server agree on location. certs_dir = str(Path(certs_dir).resolve()) + else: + # No CERTS_DIR: honour $WEIGHTSLAB_CERTS_DIR -- the directory `start --certs` + # and the backend read -- before CertAuthManager's ~/.weightslab-certs + # default. An unusable value (empty, relative) falls back with a warning. + certs_dir = env_certs_dir() # Resolve the target directory first (no filesystem work in __init__), so we # can point the generation scripts at it via WEIGHTSLAB_CERTS_DIR. manager = CertAuthManager(certs_dir=certs_dir, enable_auth=not no_auth) - exit_code = _generate_certs_with_fallback(force_certs=force_certs, certs_dir=manager.certs_dir) + exit_code = _generate_certs_with_fallback( + force_certs=force_certs, certs_dir=manager.certs_dir, force_ubuntu=force_ubuntu) if exit_code != 0: logger.error("Certificate generation failed") sys.exit(1) @@ -516,6 +553,7 @@ def ui_secure_environment(args): # Export ONLY the single source of truth for this process. os.environ["WEIGHTSLAB_CERTS_DIR"] = str(manager.certs_dir) + logger.info("====================================================================") logger.info(" Certificates generated successfully") logger.info(" gRPC auth token created") logger.info(f" Certs and token stored in: {manager.certs_dir}") @@ -527,6 +565,7 @@ def ui_secure_environment(args): "and the training backend find these certs (single source of truth):") logger.warning(f" (bash) echo 'export WEIGHTSLAB_CERTS_DIR=\"{manager.certs_dir}\"' >> ~/.bashrc && source ~/.bashrc") logger.warning(f" (Windows) setx WEIGHTSLAB_CERTS_DIR \"{manager.certs_dir}\"") + logger.info("====================================================================") # Bundled PyTorch examples, keyed by the CLI flag (e.g. --cls -> wl-classification). @@ -536,7 +575,7 @@ def ui_secure_environment(args): "seg": ("wl-segmentation", "segmentation", "PyTorch"), "det": ("wl-detection", "detection", "PyTorch"), "clus": ("wl-clustering", "clustering", "PyTorch"), - "gen": ("wl-generation", "generation", "PyTorch"), + "gen": ("wl-image-generation", "image generation", "PyTorch"), "3d_det": ("wl-3d-lidar-detection", "3D LiDAR detection", "Usecases"), "2d_det": ("wl-2d-lidar-detection", "2D LiDAR detection", "Usecases"), # One-level-at-a-time MNIST demos behind the four-way SDK approach docs. @@ -581,7 +620,7 @@ def _install_example_requirements(example_dir: Path) -> None: def example_start(args): - """`weightslab start example [--cls|--seg|--clus|--gen]`: run a bundled example. + """`weightslab start example [--cls|--seg|--det|...]`: run a bundled example. Defaults to the classification (cls) example. First installs the example's requirements (if a requirements file is present) without prompting, then runs @@ -894,7 +933,8 @@ def ui_start_native(args): server — like ``tensorboard``, the UI ships in the wheel. Unsecured HTTP by default. Pass ``--certs`` to serve HTTPS + mTLS to the - backend, derived solely from cert-file presence in $WEIGHTSLAB_CERTS_DIR. + backend, derived solely from cert-file presence in $WEIGHTSLAB_CERTS_DIR + (else ~/.weightslab-certs; see CertAuthManager.from_env_or_default). """ try: from weightslab.ui import server as ui_server @@ -939,8 +979,6 @@ def ui_start_native(args): ui_host = getattr(args, "host", None) or os.getenv("WEIGHTSLAB_UI_HOST", "0.0.0.0") preferred_ui_port, ui_port_source = _resolve_ui_port(args) - if ui_port_source == "default": - preferred_ui_port = 8080 ui_port = preferred_ui_port backend_host = (getattr(args, "backend_host", None) or os.getenv("GRPC_BACKEND_HOST", "localhost")) @@ -1018,21 +1056,23 @@ def ui_start_native(args): def _add_ui_server_flags(p: argparse.ArgumentParser) -> None: """Attach flags for the native UI server.""" p.add_argument('--port', type=int, default=None, - help='UI HTTP port (default: config ui_port, else $WL_LAST_UI_PORT, else 50051)') + help='UI HTTP port (default: ui_port from --config / $WEIGHTSLAB_EXPERIMENT_CONFIG, ' + 'else $WL_LAST_UI_PORT, else $WEIGHTSLAB_UI_PORT, else 8080). ' + 'A busy port falls back to a free one.') p.add_argument('--config', default=None, help='Experiment config file to read ui_port from (yaml/yml)') p.add_argument('--host', default=None, - help='UI bind host (default: 0.0.0.0)') + help='UI bind host (default: $WEIGHTSLAB_UI_HOST, else 0.0.0.0)') p.add_argument('--backend-host', dest='backend_host', default=None, - help='Backend gRPC host to proxy to (default: localhost)') + help='Backend gRPC host to proxy to (default: $GRPC_BACKEND_HOST, else localhost)') p.add_argument('--backend-port', dest='backend_port', type=int, default=None, - help='Backend gRPC port to proxy to (default: 50051)') + help='Backend gRPC port to proxy to (default: $GRPC_BACKEND_PORT, else 50051)') p.add_argument('--no-browser', dest='no_browser', action='store_true', help='Do not open the web browser automatically') p.add_argument('--certs', action='store_true', help='Serve HTTPS + mTLS to the backend using TLS certs from ' - '$WEIGHTSLAB_CERTS_DIR (default: unsecured HTTP). ' - 'Run `weightslab se` first to generate them.') + '$WEIGHTSLAB_CERTS_DIR, else ~/.weightslab-certs (default: ' + 'unsecured HTTP). Run `weightslab se` first to generate them.') def _add_example_kind_flags(p: argparse.ArgumentParser) -> None: @@ -1047,7 +1087,7 @@ def _add_example_kind_flags(p: argparse.ArgumentParser) -> None: group.add_argument("--clus", action="store_const", dest="example_kind", const="clus", help="Run the clustering example") group.add_argument("--gen", action="store_const", dest="example_kind", const="gen", - help="Run the generation example") + help="Run the image-generation example") group.add_argument("--3d_det", action="store_const", dest="example_kind", const="3d_det", help="Run the 3D LiDAR point-cloud detection example") group.add_argument("--2d_det", action="store_const", dest="example_kind", const="2d_det", @@ -1108,12 +1148,16 @@ def _build_parser() -> argparse.ArgumentParser: The CLI is intentionally minimal — exactly these commands: weightslab --help | -h | help - weightslab se [--force-certs] - weightslab start [--port PORT] [--config FILE] [--backend-port PORT] [--certs] - weightslab start example [--cls|--seg|--det|--clus|--gen|--3d_det|--2d_det] + weightslab se [CERTS_DIR] [--force-certs] [--force-ubuntu] + weightslab start [DIR] [--port PORT] [--config FILE] [--host HOST] + [--backend-host HOST] [--backend-port PORT] [--no-browser] [--certs] + weightslab start example [--cls|--seg|--det|--clus|--gen|--3d_det|--2d_det + |--model|--data|--config|--logger] weightslab cli [--port PORT] [--host HOST] - weightslab tunnel ENDPOINT - weightslab export --format {cvat,label_studio,v7} [OUTPUT] + weightslab tunnel [ENDPOINT] [--listen-port N] [--listen-host H] [--remote-port N] + weightslab export --format {cvat,label_studio,v7} [OUTPUT] [--origin O] + [--predictions] [--tag TAG ...] [--host H] [--port N] + weightslab agent init [--provision-only] """ parser = argparse.ArgumentParser( prog="weightslab", @@ -1123,9 +1167,12 @@ def _build_parser() -> argparse.ArgumentParser: ) sub = parser.add_subparsers(dest="command", metavar="{se,start,cli,tunnel,export,agent,help}") - # weightslab se [--force-certs] [certs_dir] + # weightslab se [--force-certs] [--force-ubuntu] [certs_dir] se_parser = sub.add_parser("se", help="Set up the secure environment (TLS certs + gRPC auth token)") se_parser.add_argument('--force-certs', action='store_true', help='Regenerate certificates even if they already exist') + se_parser.add_argument('--force-ubuntu', action='store_true', + help='Windows only: generate certs with the bash script through WSL/Ubuntu ' + 'instead of the default PowerShell script (no effect on Linux/macOS)') se_parser.add_argument('certs_dir', nargs='?', default=None, help='Custom directory for certs/token (default: $WEIGHTSLAB_CERTS_DIR or ~/.weightslab-certs)') @@ -1204,8 +1251,9 @@ def _build_parser() -> argparse.ArgumentParser: # Tolerate the swapped order: `weightslab example start [flags]` (and bare # `weightslab example`) behave exactly like `weightslab start example`. Hidden - # from --help on purpose (argparse.SUPPRESS) — a forgiving fallback. - example_alias = sub.add_parser("example", help=argparse.SUPPRESS) + # from --help on purpose — a forgiving fallback. No help= at all: argparse + # lists every subparser given one, and help=SUPPRESS prints "==SUPPRESS==". + example_alias = sub.add_parser("example") example_alias_sub = example_alias.add_subparsers(dest="example_action") example_alias_start = example_alias_sub.add_parser( "start", help="Start a bundled PyTorch example (default: classification)") diff --git a/weightslab/security/__init__.py b/weightslab/security/__init__.py index 08e8b205..f85f802c 100644 --- a/weightslab/security/__init__.py +++ b/weightslab/security/__init__.py @@ -1,5 +1,5 @@ """Security utilities for weightslab.""" -from .cert_auth_manager import CertAuthManager +from .cert_auth_manager import CertAuthManager, env_certs_dir -__all__ = ['CertAuthManager'] +__all__ = ['CertAuthManager', 'env_certs_dir'] diff --git a/weightslab/security/cert_auth_manager.py b/weightslab/security/cert_auth_manager.py index 168072c4..a8ca8804 100644 --- a/weightslab/security/cert_auth_manager.py +++ b/weightslab/security/cert_auth_manager.py @@ -41,6 +41,39 @@ def _get_user_profile() -> str: return os.environ.get('HOME') or os.path.expanduser('~') +_WARNED_CERTS_DIRS: set = set() + + +def _warn_once(message: str) -> None: + """Log a certs-dir warning once per process (the resolver runs several times).""" + if message not in _WARNED_CERTS_DIRS: + _WARNED_CERTS_DIRS.add(message) + logger.warning(message) + + +def env_certs_dir() -> Optional[str]: + """Return ``$WEIGHTSLAB_CERTS_DIR`` as a usable path, or None. + + None when the variable is unset or empty, and also when it is not an + absolute path (after ``~`` expansion and WSL/Git-Bash normalization): a + relative value resolves against each process's current directory, so the + UI, the backend and ``weightslab se`` would each look somewhere else. Such + a value is ignored with a warning and callers fall back to + ``~/.weightslab-certs``. + """ + raw = (os.environ.get('WEIGHTSLAB_CERTS_DIR') or '').strip().strip("'\"") + if not raw: + return None + value = _normalize_native_path(os.path.expanduser(raw)) + if not Path(value).is_absolute(): + _warn_once( + f"Ignoring WEIGHTSLAB_CERTS_DIR={raw!r}: not an absolute path. " + "Falling back to ~/.weightslab-certs." + ) + return None + return value + + def _generate_hex_token(byte_count: int = 32) -> str: """Generate a strong hex token using secure random.""" import secrets @@ -310,8 +343,23 @@ def initialize(self, force_certs: bool = False) -> Tuple[bool, str]: @staticmethod def from_env_or_default(enable_auth: bool = True) -> 'CertAuthManager': - """Create manager using environment variables or defaults.""" - certs_dir = os.environ.get('WEIGHTSLAB_CERTS_DIR') - if certs_dir is not None: - certs_dir = certs_dir.strip().strip("'\"") - return CertAuthManager(certs_dir=certs_dir, enable_auth=enable_auth) + """Create a manager for the certs directory to use. + + Looks in ``$WEIGHTSLAB_CERTS_DIR`` first, then ``~/.weightslab-certs``: + when the env directory has no cert set but the default one does, the + default wins, so a stale value cannot hide certs that exist. With certs + in neither, the env directory (when usable) is kept so messages point + at the directory the user chose. + """ + default = CertAuthManager(certs_dir=None, enable_auth=enable_auth) + env_dir = env_certs_dir() + if env_dir is None: + return default + manager = CertAuthManager(certs_dir=env_dir, enable_auth=enable_auth) + if manager.has_valid_certs() or not default.has_valid_certs(): + return manager + _warn_once( + f"No certs in WEIGHTSLAB_CERTS_DIR={env_dir}; using {default.certs_dir}, " + "which has them." + ) + return default From 9078eb3b87de68472d8f480c143a7cc53fb9cec2 Mon Sep 17 00:00:00 2001 From: GuillaumePELLUET Date: Thu, 1 Oct 2026 14:07:52 +0200 Subject: [PATCH 19/29] logs: terminal/file split, survive dictConfig, tqdm mirror; rewind samples on restore Logging ------- Three independent reasons the session log looked stale, all fixed: * The root logger was set to WEIGHTSLAB_LOG_LEVEL, which gates records before any handler sees them, so the file handler's DEBUG level was moot and INFO silently dropped every DEBUG record from the file too. Root now sits at min(console, file); the console handler carries the user's level and the file handler NOTSET. Cap the file with the new WEIGHTSLAB_LOG_FILE_LEVEL. * The log file was created at import and only relocated when root_log_dir came in via watch_or_edit(defaults=)/kwargs -- never for the usual case of root_log_dir inside the config dict, so a run left its whole log in %TEMP%. Relocation now happens once, right after _resolve_configured_root_log_dir, into /weightslab_logs/ (the layout setup_logging already uses, so the file never moves between two conventions mid-run). * IPKernelApp.initialize() -> traitlets -> logging.config.dictConfig -> _clearExistingHandlers -> logging.shutdown() closes EVERY handler in the process without detaching them from the root logger. StreamHandler survives; FileHandler does not, and CPython refuses to reopen a closed mode='w' handler (bpo-42378), so the log died mid-run with no error while the terminal carried on. The session file now opens in append mode (filenames are timestamped, so appending == truncating for a fresh file) and ensure_logging_intact() repairs the root logger; the notebook kernel calls it right after initialize(). Also: the exit notice used this module's own print shim, so it went through logging during interpreter teardown and surfaced as "I/O operation on closed file" at the end of every run. It uses builtins.print now. tqdm -> log mirror ------------------ A tqdm bar paints onto a terminal and never touches logging, so the log file had no record of a run's own progress. tqdm_logging samples live bars on a timer and writes one compact line per bar that actually moved, to a weightslab.progress channel the terminal handler filters out (the live bar is already there; the terminal copy goes through tqdm.write so the bar redraws cleanly). WEIGHTSLAB_TQDM_LOG_INTERVAL (default 30s, 0 disables), WEIGHTSLAB_TQDM_LOG_TO_TERMINAL to turn the echo off. Per-sample rewind on checkpoint restore --------------------------------------- Restoring weights moves the model's age backwards, but the ledger kept whatever the later steps wrote, so the grid showed a sample's loss from step 900 beside a model back at 400. * LoggerQueue.get_per_sample_state_at_step: last value per signal at or before a step, plus last_seen = MAX(step) and nb_seen = COUNT(DISTINCT step). Includes the run's _ evaluation hashes, since eval passes bump the same counters. * DataFrameManager.rewind_to_step: rewrites only samples whose last_seen is ahead of the step -- signals from history (NaN where a signal has nothing that old), counters recomputed, prediction/prediction_raw cleared, targets kept. * CheckpointManager._rewind_sample_state, called at the end of load_state. Runs only when the loaded hash is the one already loaded (step axes are comparable only within an experiment) and skips with a warning when the history is empty, so a not-yet-loaded logger cannot wipe the ledger. Embedded notebook kernel startup order -------------------------------------- _install_thread_routed_streams now runs immediately after initialize() instead of ~2s later: anything capturing sys.stdout/stderr in that gap was bound to the unrouted OutStream for good, which is why the trainer's tqdm bar disappeared from the terminal once the notebook kernel was enabled. flush_interval is set on the raw OutStreams before the wrap -- after it, the assignment lands on the wrapper and cell output stops streaming live. Both orderings are now pinned by tests. wl-classification example ------------------------- It hardcoded skip_previous_auto_load=True, so a restart never resumed from the newest weights and config.yaml's documented skip_checkpoint_load was inert for the model (the kwarg is only ever OR-ed with the config, so a True kwarg can never be turned off). It reads the config now. The config comment said the opposite of what the flag does. .gitignore ---------- The bare `data` rule excluded the tests/data/ and weightslab/data/ directories, and the intended escape hatches never worked: a leading ./ does not match, since patterns are relative to the .gitignore. Consequences: 13 test files under tests/data/ were invisible to git and CI, and weightslab/data/h5_recovery.py -- imported at module load by dataframe_manager.py and checkpoint_manager.py -- was never committed, so a fresh clone could not import weightslab. Both rules fixed and the missing files added. Co-Authored-By: Claude Opus 5 --- .gitignore | 4 +- AGENTS.md | 21 +- docs/configuration.rst | 39 +- docs/usage/parameters.rst | 14 +- docs/weights_studio/configuration.rst | 6 +- docs/weights_studio_ui/more/configuration.rst | 6 +- tests/backend/test_logger_per_sample_state.py | 144 ++++ tests/components/test_checkpoint_rewind.py | 137 +++ tests/data/__init__.py | 0 tests/data/test_boolean_tag_registry.py | 83 ++ tests/data/test_categorical_tags.py | 168 ++++ tests/data/test_data_samples_with_ops.py | 783 ++++++++++++++++++ tests/data/test_data_service_metadata_copy.py | 97 +++ tests/data/test_data_utils_unit.py | 117 +++ tests/data/test_dataframe_data_invariants.py | 243 ++++++ tests/data/test_dataframe_manager_unit.py | 541 ++++++++++++ tests/data/test_dataframe_rewind.py | 173 ++++ tests/data/test_flush_pipeline.py | 276 ++++++ tests/data/test_h5_array_store.py | 216 +++++ tests/data/test_h5_dataframe_store.py | 243 ++++++ tests/data/test_point_cloud_utils.py | 330 ++++++++ tests/data/test_save_signals_e2e.py | 213 +++++ .../services/test_notebook_service_unit.py | 48 ++ tests/utils/test_logs_unit.py | 353 +++++++- tests/utils/test_tqdm_logging_unit.py | 231 ++++++ weightslab/AGENTS.md | 18 +- weightslab/__init__.py | 14 +- weightslab/backend/logger.py | 84 ++ weightslab/components/checkpoint_manager.py | 88 +- weightslab/data/array_proxy.py | 13 + weightslab/data/dataframe_manager.py | 173 ++++ weightslab/data/h5_array_store.py | 193 ++++- weightslab/data/h5_dataframe_store.py | 52 +- weightslab/data/h5_recovery.py | 149 ++++ .../PyTorch/wl-classification/config.yaml | 4 +- .../PyTorch/wl-classification/main.py | 14 +- weightslab/src.py | 112 ++- .../trainer/services/notebook_service.py | 38 +- weightslab/utils/logs.py | 347 ++++++-- weightslab/utils/tools.py | 40 + weightslab/utils/tqdm_logging.py | 216 +++++ 41 files changed, 5818 insertions(+), 223 deletions(-) create mode 100644 tests/backend/test_logger_per_sample_state.py create mode 100644 tests/components/test_checkpoint_rewind.py create mode 100644 tests/data/__init__.py create mode 100644 tests/data/test_boolean_tag_registry.py create mode 100644 tests/data/test_categorical_tags.py create mode 100644 tests/data/test_data_samples_with_ops.py create mode 100644 tests/data/test_data_service_metadata_copy.py create mode 100644 tests/data/test_data_utils_unit.py create mode 100644 tests/data/test_dataframe_data_invariants.py create mode 100644 tests/data/test_dataframe_manager_unit.py create mode 100644 tests/data/test_dataframe_rewind.py create mode 100644 tests/data/test_flush_pipeline.py create mode 100644 tests/data/test_h5_array_store.py create mode 100644 tests/data/test_h5_dataframe_store.py create mode 100644 tests/data/test_point_cloud_utils.py create mode 100644 tests/data/test_save_signals_e2e.py create mode 100644 tests/utils/test_tqdm_logging_unit.py create mode 100644 weightslab/data/h5_recovery.py create mode 100644 weightslab/utils/tqdm_logging.py diff --git a/.gitignore b/.gitignore index 2589168b..9999cf20 100644 --- a/.gitignore +++ b/.gitignore @@ -19,8 +19,8 @@ venv runs data outputs -!./tests/data/ -!./weightslab/data/ +!tests/data/ +!weightslab/data/ MagicMock drop htmlcov diff --git a/AGENTS.md b/AGENTS.md index 5786cc5d..a6483ff2 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -100,8 +100,9 @@ Working starting points live in (each is a `main.py` + `config.yaml`) — find the closest example and mirror it. UI deployment details (port, TLS, certs) are documented in -`weightslab/docs/weights_studio.rst`. TLS is opt-in: run `weightslab se` once, -then `weightslab start --certs`. On Windows `se` uses the PowerShell script and +`weightslab/docs/weights_studio.rst`. TLS turns on once certs exist: run +`weightslab se` once, and `weightslab start` and the backend then use them +automatically (`weightslab start --no-certs` forces HTTP). On Windows `se` uses the PowerShell script and the Windows `openssl`; `weightslab se --force-ubuntu` uses the bash script through WSL instead. @@ -177,10 +178,12 @@ ones when debugging: | Variable | Default | Why you touch it | |---|---|---| -| `WEIGHTSLAB_LOG_LEVEL` | `INFO` | Set `DEBUG` to see what's happening. (`WATCHDOG` level sits between WARNING/ERROR.) | +| `WEIGHTSLAB_LOG_LEVEL` | `INFO` | **Terminal only.** Set `DEBUG` to see what's happening. (`WATCHDOG` level sits between WARNING/ERROR.) | +| `WEIGHTSLAB_LOG_FILE_LEVEL` | *(unset = all)* | The session log file keeps every record whatever the terminal shows; set this to cap the file too. Log lives in `/weightslab_logs/`. | +| `WEIGHTSLAB_TQDM_LOG_INTERVAL` | `30` | Seconds between snapshots of live `tqdm` bars into the log (`0` disables). tqdm never goes through `logging`, so without this the file has no record of training progress. | | `GRPC_BACKEND_HOST` / `GRPC_BACKEND_PORT` | `0.0.0.0` / `50051` | Backend gRPC bind address. | -| `GRPC_TLS_ENABLED` | `0` | TLS on the gRPC socket. Set `1` with `weightslab start --certs`. | -| `GRPC_TLS_REQUIRE_CLIENT_AUTH` | `0` | mTLS. Must match what `weightslab start --certs` presents. | +| `GRPC_TLS_ENABLED` | `0` | TLS on the gRPC socket. Set to `1` automatically when certs are found; `0`/`false` forces plaintext (also for `weightslab start`). | +| `GRPC_TLS_REQUIRE_CLIENT_AUTH` | `0` | mTLS. Must match what `weightslab start` presents when it uses certs. | | `WEIGHTSLAB_CERTS_DIR` | `~/.weightslab-certs` | Where cert files are looked up (single source of truth). Falls back to `~/.weightslab-certs` when unset, not an absolute path, or holding no certs. | | `GRPC_AUTH_TOKEN` | *(unset)* | Optional metadata-token auth on top of mTLS. | | `GRPC_MAX_MESSAGE_BYTES` | `268435456` (256 MB) | Raise it if large tensors/image batches fail. | @@ -218,9 +221,11 @@ distilled from issues hit in development). **UI loads but the sample grid is empty / "failed to fetch" / gRPC errors.** The wire path (§1) is broken somewhere. Check in order: (1) backend actually serving on `0.0.0.0:50051`; (2) `weightslab start` is running and the browser -can reach it on `:8080`; (3) **TLS mismatch** if using `--certs` — run -`weightslab se` first and export `WEIGHTSLAB_CERTS_DIR`. For local debugging -drop TLS entirely (omit `--certs`; `GRPC_TLS_ENABLED=0`). +can reach it on `:8080`; (3) **TLS mismatch** — the UI and the backend each +turn TLS on when they find certs, so both must see the same +`WEIGHTSLAB_CERTS_DIR` (the browser console prints `TLS: ENABLED/DISABLED`). +For local debugging drop TLS on both sides (`weightslab start --no-certs`; +`GRPC_TLS_ENABLED=0` for the backend). **`weightslab se` hangs with no output (Windows).** Only the WSL path can do this: `--force-ubuntu`, or the fallback after the diff --git a/docs/configuration.rst b/docs/configuration.rst index 15f5d3f2..0ed5610c 100644 --- a/docs/configuration.rst +++ b/docs/configuration.rst @@ -410,18 +410,43 @@ Logging - Description * - ``WEIGHTSLAB_LOG_LEVEL`` - ``INFO`` - - Log level for all WeightsLab Python components. + - Minimum level printed **to the terminal**. The session log file is not + affected: it records everything regardless of this setting (see + ``WEIGHTSLAB_LOG_FILE_LEVEL``), so a quiet terminal still leaves a + full-fidelity log on disk. Accepted values: ``DEBUG``, ``INFO``, ``WARNING``, ``ERROR``, ``WATCHDOG``. ``WATCHDOG`` (level 35) sits between WARNING and ERROR and is used for watchdog/restart events. - * - ``WEIGHTSLAB_LOG_TO_FILE`` + * - ``WEIGHTSLAB_LOG_FILE_LEVEL`` + - *(unset — everything)* + - Minimum level written to the session log file. Unset means no + restriction, which is the point: the terminal is filtered, the file is + complete. Set it (e.g. ``INFO``) to cap the file too when the full log + is more than you want on disk. + * - ``WEIGHTSLAB_TQDM_LOG_INTERVAL`` + - ``30`` + - Seconds between snapshots of any live ``tqdm`` progress bar into the + session log (``0`` disables). A bar paints itself onto the terminal and + never goes through ``logging``, so without this the log file has no + record of the run's own progress. Lines are sampled, not streamed: an + unchanged bar is not repeated. + * - ``WEIGHTSLAB_TQDM_LOG_TO_TERMINAL`` - ``0`` - - Write logs to a rotating file in addition to stdout. - Set to ``1`` to enable. + - By default the sampled progress lines go to the file only, since the + live bar is already on the terminal. Set to ``1`` to print them too. + * - ``WEIGHTSLAB_LOG_TO_FILE`` + - ``true`` + - Write a session log file in addition to stdout. Set to ``false`` to + disable. Only the main process writes one — ``DataLoader`` workers and + other spawned children keep a terminal-only logger. * - ``WEIGHTSLAB_ROOT_LOG_DIR`` - - *(training script dir)* - - Root directory where training log snapshots are saved. - Defaults to a ``root_log_dir/`` folder next to your training script. + - *(temporary directory)* + - Experiment directory (checkpoints, reports, notebook, logs). Also what + ``weightslab start [DIR]`` exports. The session log starts in + ``/weightslab_logs/`` and, when it is unset, starts in a temporary + directory and is **moved** into ``/weightslab_logs/`` as + soon as the experiment's ``root_log_dir`` resolves — so the log always + ends up beside the checkpoints. The path is printed when the process exits. * - ``AUDIT_LOG_FORMAT`` - ``json`` - Output format for audit logs tracking all user interactions through gRPC. diff --git a/docs/usage/parameters.rst b/docs/usage/parameters.rst index 5a4df2c3..cdecad6d 100644 --- a/docs/usage/parameters.rst +++ b/docs/usage/parameters.rst @@ -288,14 +288,20 @@ Logging & debug - Description * - ``WEIGHTSLAB_LOG_LEVEL`` - ``INFO`` - - Log verbosity for all WeightsLab Python components. + - Minimum level printed **to the terminal**; the session log file + keeps everything regardless. Accepted: ``DEBUG``, ``INFO``, ``WARNING``, ``ERROR``. ``WATCHDOG`` (level 35) is a custom level between WARNING and ERROR reserved for watchdog/restart events. + * - ``WEIGHTSLAB_LOG_FILE_LEVEL`` + - *(unset — everything)* + - Minimum level written to the session log file. Set it to cap + the file as well as the terminal. * - ``WEIGHTSLAB_LOG_TO_FILE`` - - ``0`` - - Set to ``1`` to write logs to a rotating file in the system - temp directory in addition to stdout. + - ``true`` + - Set to ``false`` to skip the session log file. The file lives in + ``/weightslab_logs/`` once the experiment directory + resolves; its path is printed when the process exits. * - ``WEIGHTSLAB_SUPPRESS_BANNER`` - ``0`` - Set to ``1`` to suppress the ASCII art startup banner. diff --git a/docs/weights_studio/configuration.rst b/docs/weights_studio/configuration.rst index 3e9b7dd7..e13cd663 100644 --- a/docs/weights_studio/configuration.rst +++ b/docs/weights_studio/configuration.rst @@ -7,7 +7,9 @@ Backend environment variables (set before starting ``wl.serve()``) +----------------------------------+-------------------------+----------------------------------------------------+ | Variable | Default | Description | +==================================+=========================+====================================================+ -| ``WEIGHTSLAB_LOG_LEVEL`` | ``INFO`` | Log level (``DEBUG``, ``INFO``, ...) | +| ``WEIGHTSLAB_LOG_LEVEL`` | ``INFO`` | Terminal log level (``DEBUG``, ``INFO``, ...) | ++----------------------------------+-------------------------+----------------------------------------------------+ +| ``WEIGHTSLAB_LOG_FILE_LEVEL`` | *(unset)* | Log file level; unset keeps every record | +----------------------------------+-------------------------+----------------------------------------------------+ | ``GRPC_BACKEND_HOST`` | ``0.0.0.0`` | Host the backend gRPC server binds to | +----------------------------------+-------------------------+----------------------------------------------------+ @@ -41,7 +43,7 @@ UI server environment variables (set before ``weightslab start``) +---------------------------+-------------------------+--------------------------------------------------+ | ``GRPC_BACKEND_PORT`` | ``50051`` | Backend gRPC port to proxy to | +---------------------------+-------------------------+--------------------------------------------------+ -| ``WEIGHTSLAB_CERTS_DIR`` | ``~/.weightslab-certs`` | Certs dir (read when ``--certs``); when it | +| ``WEIGHTSLAB_CERTS_DIR`` | ``~/.weightslab-certs`` | Certs dir; HTTPS when it has certs; when it | | | | holds none, ``~/.weightslab-certs`` is used | +---------------------------+-------------------------+--------------------------------------------------+ | ``WEIGHTSLAB_OPENCODE_PORT`` | ``4096`` | Port the agent (OpenCode) server is started on; | diff --git a/docs/weights_studio_ui/more/configuration.rst b/docs/weights_studio_ui/more/configuration.rst index 3e9b7dd7..e13cd663 100644 --- a/docs/weights_studio_ui/more/configuration.rst +++ b/docs/weights_studio_ui/more/configuration.rst @@ -7,7 +7,9 @@ Backend environment variables (set before starting ``wl.serve()``) +----------------------------------+-------------------------+----------------------------------------------------+ | Variable | Default | Description | +==================================+=========================+====================================================+ -| ``WEIGHTSLAB_LOG_LEVEL`` | ``INFO`` | Log level (``DEBUG``, ``INFO``, ...) | +| ``WEIGHTSLAB_LOG_LEVEL`` | ``INFO`` | Terminal log level (``DEBUG``, ``INFO``, ...) | ++----------------------------------+-------------------------+----------------------------------------------------+ +| ``WEIGHTSLAB_LOG_FILE_LEVEL`` | *(unset)* | Log file level; unset keeps every record | +----------------------------------+-------------------------+----------------------------------------------------+ | ``GRPC_BACKEND_HOST`` | ``0.0.0.0`` | Host the backend gRPC server binds to | +----------------------------------+-------------------------+----------------------------------------------------+ @@ -41,7 +43,7 @@ UI server environment variables (set before ``weightslab start``) +---------------------------+-------------------------+--------------------------------------------------+ | ``GRPC_BACKEND_PORT`` | ``50051`` | Backend gRPC port to proxy to | +---------------------------+-------------------------+--------------------------------------------------+ -| ``WEIGHTSLAB_CERTS_DIR`` | ``~/.weightslab-certs`` | Certs dir (read when ``--certs``); when it | +| ``WEIGHTSLAB_CERTS_DIR`` | ``~/.weightslab-certs`` | Certs dir; HTTPS when it has certs; when it | | | | holds none, ``~/.weightslab-certs`` is used | +---------------------------+-------------------------+--------------------------------------------------+ | ``WEIGHTSLAB_OPENCODE_PORT`` | ``4096`` | Port the agent (OpenCode) server is started on; | diff --git a/tests/backend/test_logger_per_sample_state.py b/tests/backend/test_logger_per_sample_state.py new file mode 100644 index 00000000..5e7065de --- /dev/null +++ b/tests/backend/test_logger_per_sample_state.py @@ -0,0 +1,144 @@ +"""Tests for ``LoggerQueue.get_per_sample_state_at_step``. + +The query behind the post-restore rewind: what every sample looked like when +the model was a given age — the last value each signal held at or before that +step, plus the seen-counters implied by the same rows. +""" + +import math +import unittest +from unittest.mock import patch + +from weightslab.backend.logger import LoggerQueue + + +EXP = "hash_of_the_experiment" +OTHER_EXP = "hash_of_another_run" + + +def _lg(exp_hash=EXP) -> LoggerQueue: + """Unregistered LoggerQueue on a private in-memory DB. + + ``register=False`` is not enough for isolation: ``__init__`` looks up the + ledger's checkpoint manager regardless and, if one is registered, rebinds + the connection from ``:memory:`` to that experiment's on-disk + ``loggers.duckdb``. Any earlier suite that leaves a manager in the ledger + therefore makes every instance here share one database, and the rows from + each test pile up into the next one's counts. Patching the lookup keeps the + connection in memory and private to this instance. + """ + with patch("weightslab.backend.logger.get_checkpoint_manager", return_value=None): + lg = LoggerQueue(register=False) + lg.chkpt_manager = type( + "_FakeCM", (), {"get_current_experiment_hash": staticmethod(lambda: exp_hash)})() + return lg + + +def _add(lg, signal, sample_id, step, value): + lg.add_scalars(signal, {signal: value}, step, + signal_per_sample={sample_id: value}, aggregate_by_step=False) + + +class TestPerSampleStateAtStep(unittest.TestCase): + def setUp(self): + self.lg = _lg() + for step, value in ((1, 0.9), (2, 0.7), (5, 0.4), (9, 0.1)): + _add(self.lg, "loss", "1", step, value) + for step, value in ((1, 0.8), (5, 0.3)): + _add(self.lg, "loss", "2", step, value) + # Only ever recorded after the step we rewind to. + _add(self.lg, "late", "1", 9, 42.0) + + # Guards the isolation _lg() buys: a shared database would carry rows + # from earlier tests in here and silently inflate every nb_seen below. + self.lg._flush_stage() + self.assertEqual( + self.lg._conn.execute("SELECT count(*) FROM per_sample").fetchone()[0], 7, + "per_sample is not isolated to this test") + + def tearDown(self): + self.lg.stop_background_flush() + + def test_returns_the_last_value_at_or_before_the_step(self): + state = self.lg.get_per_sample_state_at_step(5, exp_hash=EXP) + + self.assertAlmostEqual(state["1"]["signals"]["loss"], 0.4, places=6) + self.assertAlmostEqual(state["2"]["signals"]["loss"], 0.3, places=6) + + def test_ignores_signals_first_recorded_after_the_step(self): + state = self.lg.get_per_sample_state_at_step(5, exp_hash=EXP) + self.assertNotIn("late", state["1"]["signals"]) + + # ...but keeps it once the step is late enough to include it. + later = self.lg.get_per_sample_state_at_step(9, exp_hash=EXP) + self.assertAlmostEqual(later["1"]["signals"]["late"], 42.0, places=6) + + def test_counters_are_derived_from_the_same_rows(self): + state = self.lg.get_per_sample_state_at_step(5, exp_hash=EXP) + + self.assertEqual(state["1"]["last_seen"], 5) + self.assertEqual(state["1"]["nb_seen"], 3) # steps 1, 2, 5 + self.assertEqual(state["2"]["last_seen"], 5) + self.assertEqual(state["2"]["nb_seen"], 2) # steps 1, 5 + + def test_nb_seen_counts_distinct_steps_not_rows(self): + # Two signals at one step is one sighting of the sample. + _add(self.lg, "second_signal", "2", 5, 1.0) + state = self.lg.get_per_sample_state_at_step(5, exp_hash=EXP) + self.assertEqual(state["2"]["nb_seen"], 2) + + def test_samples_with_no_history_that_old_are_absent(self): + _add(self.lg, "loss", "newcomer", 8, 0.5) + + state = self.lg.get_per_sample_state_at_step(5, exp_hash=EXP) + self.assertNotIn("newcomer", state) + self.assertIn("newcomer", self.lg.get_per_sample_state_at_step(8, exp_hash=EXP)) + + def test_step_before_any_history_returns_nothing(self): + self.assertEqual(self.lg.get_per_sample_state_at_step(0, exp_hash=EXP), {}) + + def test_other_experiments_are_excluded(self): + # Same sample id, same step range, different run. + self.lg.ingest_per_sample("loss", OTHER_EXP, [("1", 3, 0.123)]) + + state = self.lg.get_per_sample_state_at_step(5, exp_hash=EXP) + self.assertAlmostEqual(state["1"]["signals"]["loss"], 0.4, places=6) + + everything = self.lg.get_per_sample_state_at_step(5, exp_hash=None) + self.assertEqual(everything["1"]["nb_seen"], 4) # steps 1, 2, 3, 5 + + def test_evaluation_hashes_count_towards_the_experiment(self): + eval_hash = f"{EXP}_1" + self.lg.start_evaluation_mode("test", eval_hash) + _add(self.lg, "loss", "1", 4, 0.55) + self.lg.stop_evaluation_mode(model_age=4) + + with_evals = self.lg.get_per_sample_state_at_step(5, exp_hash=EXP) + self.assertEqual(with_evals["1"]["nb_seen"], 4) # steps 1, 2, 4 (eval), 5 + + without = self.lg.get_per_sample_state_at_step( + 5, exp_hash=EXP, include_evaluations=False) + self.assertEqual(without["1"]["nb_seen"], 3) + + def test_metric_names_filters_signals_but_not_counters(self): + state = self.lg.get_per_sample_state_at_step(9, exp_hash=EXP, metric_names=["loss"]) + + self.assertEqual(set(state["1"]["signals"]), {"loss"}) + # "late" still marks step 9 as a sighting even though it was filtered out. + self.assertEqual(state["1"]["last_seen"], 9) + self.assertEqual(state["1"]["nb_seen"], 4) + + def test_nan_values_survive_as_nan(self): + self.lg.ingest_per_sample("loss", EXP, [("nan_sample", 3, float("nan"))]) + state = self.lg.get_per_sample_state_at_step(5, exp_hash=EXP) + self.assertTrue(math.isnan(state["nan_sample"]["signals"]["loss"])) + + def test_reads_values_still_in_the_staging_buffer(self): + # add_scalars only stages; the query must flush before reading. + _add(self.lg, "loss", "3", 2, 0.66) + state = self.lg.get_per_sample_state_at_step(5, exp_hash=EXP) + self.assertAlmostEqual(state["3"]["signals"]["loss"], 0.66, places=6) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/components/test_checkpoint_rewind.py b/tests/components/test_checkpoint_rewind.py new file mode 100644 index 00000000..13cc9d11 --- /dev/null +++ b/tests/components/test_checkpoint_rewind.py @@ -0,0 +1,137 @@ +"""Tests for ``CheckpointManager._rewind_sample_state``. + +The orchestration between the signal history and the ledger after a restore: +when it runs, when it deliberately does not, and that it never takes a restore +down with it. +""" + +import tempfile +import unittest +from unittest.mock import MagicMock, patch + +import numpy as np +import pandas as pd + +from weightslab.backend.logger import LoggerQueue +from weightslab.components.checkpoint_manager import CheckpointManager +from weightslab.data.dataframe_manager import LedgeredDataFrameManager +from weightslab.data.sample_stats import SampleStats + + +EXP = "experiment_hash_aaaa" +OTHER_EXP = "experiment_hash_bbbb" + +LAST_SEEN = SampleStats.Ex.LAST_SEEN.value +NB_SEEN = SampleStats.Ex.NB_SEEN.value +PREDICTION = SampleStats.Ex.PREDICTION.value + + +class RewindOrchestrationTestBase(unittest.TestCase): + def setUp(self): + self._tmpdir = tempfile.TemporaryDirectory() + self.addCleanup(self._tmpdir.cleanup) + self.cm = CheckpointManager(root_log_dir=self._tmpdir.name) + self.cm.current_exp_hash = EXP + + def _patch_ledger(self, dataframe, logger_queue): + patcher = patch("weightslab.components.checkpoint_manager.ledgers") + ledgers = patcher.start() + self.addCleanup(patcher.stop) + ledgers.get_dataframe.return_value = dataframe + ledgers.get_logger.return_value = logger_queue + return ledgers + + +class TestRewindGuards(RewindOrchestrationTestBase): + def setUp(self): + super().setUp() + self.dfm = MagicMock() + self.dfm.rewind_to_step.return_value = 3 + self.lg = MagicMock() + self.lg.get_per_sample_state_at_step.return_value = {"1": {"signals": {}, "last_seen": 5, "nb_seen": 2}} + self._patch_ledger(self.dfm, self.lg) + + def test_rewinds_when_reloading_the_current_experiment(self): + self.assertEqual(self.cm._rewind_sample_state(EXP, 5), 3) + + self.lg.get_per_sample_state_at_step.assert_called_once_with(5, exp_hash=EXP) + step, state = self.dfm.rewind_to_step.call_args.args + self.assertEqual(step, 5) + self.assertIn("1", state) + + def test_no_restored_step_means_nothing_to_rewind_to(self): + self.assertEqual(self.cm._rewind_sample_state(EXP, None), 0) + self.dfm.rewind_to_step.assert_not_called() + + def test_step_zero_still_rewinds(self): + # Restoring the initial checkpoint is a real rewind, not "no step". + self.cm._rewind_sample_state(EXP, 0) + self.lg.get_per_sample_state_at_step.assert_called_once_with(0, exp_hash=EXP) + + def test_a_different_experiment_is_left_alone(self): + # Step numbers are only comparable within one experiment; that run's + # per-sample state comes from its own snapshot instead. + self.assertEqual(self.cm._rewind_sample_state(OTHER_EXP, 5), 0) + self.dfm.rewind_to_step.assert_not_called() + + def test_empty_history_does_not_wipe_the_ledger(self): + self.lg.get_per_sample_state_at_step.return_value = {} + + self.assertEqual(self.cm._rewind_sample_state(EXP, 5), 0) + self.dfm.rewind_to_step.assert_not_called() + + def test_a_logger_without_the_query_is_skipped(self): + self._patch_ledger(self.dfm, object()) + self.assertEqual(self.cm._rewind_sample_state(EXP, 5), 0) + self.dfm.rewind_to_step.assert_not_called() + + def test_a_failure_never_fails_the_restore(self): + self.lg.get_per_sample_state_at_step.side_effect = RuntimeError("db gone") + self.assertEqual(self.cm._rewind_sample_state(EXP, 5), 0) + + +class TestRewindEndToEnd(RewindOrchestrationTestBase): + """Real LoggerQueue + real ledger, so the two halves are checked together.""" + + def setUp(self): + super().setUp() + self.lg = LoggerQueue(register=False) + self.lg.chkpt_manager = type( + "_FakeCM", (), {"get_current_experiment_hash": staticmethod(lambda: EXP)})() + self.addCleanup(self.lg.stop_background_flush) + + for step, value in ((1, 0.9), (5, 0.4), (9, 0.1)): + self.lg.add_scalars("loss", {"loss": value}, step, + signal_per_sample={"1": value}, aggregate_by_step=False) + + self.dfm = LedgeredDataFrameManager( + enable_flushing_threads=False, enable_h5_persistence=False) + self.dfm.upsert_df( + pd.DataFrame([{ + "sample_id": "1", "origin": "train", "signals//loss": 0.1, + LAST_SEEN: 9, NB_SEEN: 3, PREDICTION: np.array([1, 2]), + }]).set_index("sample_id"), origin="train") + + self._patch_ledger(self.dfm, self.lg) + + def _cell(self, column): + return self.dfm.get_df_view().loc[("1", 0), column] + + def test_restoring_step_five_puts_the_sample_back_where_it_was(self): + self.assertEqual(self.cm._rewind_sample_state(EXP, 5), 1) + + self.assertAlmostEqual(self._cell("signals//loss"), 0.4, places=6) + self.assertEqual(int(self._cell(LAST_SEEN)), 5) + self.assertEqual(int(self._cell(NB_SEEN)), 2) # steps 1 and 5 + self.assertIsNone(self._cell(PREDICTION)) + + def test_restoring_the_latest_step_changes_nothing(self): + self.assertEqual(self.cm._rewind_sample_state(EXP, 9), 0) + + self.assertAlmostEqual(self._cell("signals//loss"), 0.1, places=6) + self.assertEqual(int(self._cell(LAST_SEEN)), 9) + self.assertIsNotNone(self._cell(PREDICTION)) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/data/__init__.py b/tests/data/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/tests/data/test_boolean_tag_registry.py b/tests/data/test_boolean_tag_registry.py new file mode 100644 index 00000000..7514f03f --- /dev/null +++ b/tests/data/test_boolean_tag_registry.py @@ -0,0 +1,83 @@ +"""Boolean tags declared before any sample wears them. + +The UI's painter lets a tag be created and then painted. Between those two +moments the tag exists nowhere in the data -- and the tag list the UI shows is +read back out of the dataframe's ``tag:`` columns, so without a registry +the freshly created tag disappears on the next refresh. These tests cover the +declaration itself; ``DataService._get_unique_tags`` reporting it is covered in +tests/trainer/services/test_trainer_services_unit.py. +""" + +import unittest + +import pandas as pd + +from weightslab.data.dataframe_manager import LedgeredDataFrameManager + + +class TestBooleanTagRegistry(unittest.TestCase): + def _mgr(self): + return LedgeredDataFrameManager( + enable_flushing_threads=False, enable_h5_persistence=False + ) + + def _with_samples(self): + mgr = self._mgr() + df = pd.DataFrame( + {"sample_id": [1, 2], "origin": "train", "tag:painted": [True, False]} + ).set_index("sample_id") + mgr.upsert_df(df, origin="train") + return mgr + + def test_declares_tag_and_creates_its_column_unset(self): + mgr = self._with_samples() + self.assertTrue(mgr.register_boolean_tag("fresh")) + self.assertIn("fresh", mgr.get_declared_boolean_tags()) + + combined = mgr.get_combined_df() + self.assertIn("tag:fresh", combined.columns) + # Declared, not applied: no sample wears it. + self.assertFalse(combined["tag:fresh"].any()) + + def test_tolerates_the_tag_prefix_and_rejects_empty_names(self): + mgr = self._with_samples() + self.assertTrue(mgr.register_boolean_tag("tag:prefixed")) + self.assertIn("prefixed", mgr.get_declared_boolean_tags()) + self.assertNotIn("tag:prefixed", mgr.get_declared_boolean_tags()) + + self.assertFalse(mgr.register_boolean_tag(" ")) + self.assertFalse(mgr.register_boolean_tag("None")) + + def test_registering_twice_is_idempotent_and_keeps_painted_values(self): + mgr = self._with_samples() + self.assertTrue(mgr.register_boolean_tag("painted")) + self.assertTrue(mgr.register_boolean_tag("painted")) + self.assertEqual( + [t for t in mgr.get_declared_boolean_tags() if t == "painted"], ["painted"] + ) + # The existing column is left alone -- re-declaring a tag in use must not + # wipe the samples that already carry it. + self.assertTrue(mgr.get_combined_df()["tag:painted"].any()) + + def test_never_shadows_a_categorical_tag(self): + mgr = self._with_samples() + mgr.register_categorical_tag("weather", ["rainy", "sunny"]) + self.assertFalse(mgr.register_boolean_tag("weather")) + self.assertNotIn("weather", mgr.get_declared_boolean_tags()) + + def test_declared_without_any_samples_yet(self): + # No dataset registered: there is no row to hang a column on, but the tag + # is still remembered so it shows up once data arrives. + mgr = self._mgr() + self.assertTrue(mgr.register_boolean_tag("early")) + self.assertEqual(mgr.get_declared_boolean_tags(), ["early"]) + + def test_unregister_forgets_the_tag(self): + mgr = self._with_samples() + mgr.register_boolean_tag("doomed") + mgr.unregister_boolean_tag("tag:doomed") + self.assertNotIn("doomed", mgr.get_declared_boolean_tags()) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/data/test_categorical_tags.py b/tests/data/test_categorical_tags.py new file mode 100644 index 00000000..9936b87d --- /dev/null +++ b/tests/data/test_categorical_tags.py @@ -0,0 +1,168 @@ +import tempfile +import unittest +from pathlib import Path + +import numpy as np +import pandas as pd + +from weightslab.data.dataframe_manager import LedgeredDataFrameManager +from weightslab.data.h5_dataframe_store import H5DataFrameStore + + +class TestCategoricalTagRegistryManager(unittest.TestCase): + """Registry behaviour on LedgeredDataFrameManager (no H5 persistence).""" + + def _mgr(self): + return LedgeredDataFrameManager(enable_flushing_threads=False, enable_h5_persistence=False) + + def test_register_merge_and_replace(self): + mgr = self._mgr() + self.assertEqual(mgr.register_categorical_tag("weather", ["rainy", "sunny"]), ["rainy", "sunny"]) + # Merge keeps order, dedups + self.assertEqual(mgr.register_categorical_tag("weather", ["sunny", "cloudy"]), ["rainy", "sunny", "cloudy"]) + # Replace wipes previous + self.assertEqual(mgr.register_categorical_tag("weather", ["fog"], replace=True), ["fog"]) + self.assertTrue(mgr.is_categorical_tag("weather")) + self.assertTrue(mgr.is_categorical_tag("tag:weather")) # prefix tolerated + self.assertFalse(mgr.is_categorical_tag("does_not_exist")) + + def test_register_strips_prefix_and_cleans(self): + mgr = self._mgr() + # "tag:" prefix on the name is stripped; empty/None/nan categories dropped + out = mgr.register_categorical_tag("tag:quality", ["high", "", None, "low", "nan"]) + self.assertEqual(out, ["high", "low"]) + self.assertIn("quality", mgr.get_categorical_tags()) + + def test_auto_detect_from_string_tag_column(self): + mgr = self._mgr() + df = pd.DataFrame( + {"sample_id": [1, 2, 3], "origin": "train", "tag:weather": ["rainy", "sunny", "rainy"]} + ).set_index("sample_id") + mgr.upsert_df(df, origin="train") + reg = mgr.get_categorical_tags() + self.assertIn("weather", reg) + self.assertEqual(set(reg["weather"]), {"rainy", "sunny"}) + + def test_boolean_tag_not_registered_as_categorical(self): + mgr = self._mgr() + df = pd.DataFrame( + {"sample_id": [1, 2], "origin": "train", "tag:is_urban": [True, False]} + ).set_index("sample_id") + mgr.upsert_df(df, origin="train") + self.assertNotIn("is_urban", mgr.get_categorical_tags()) + + def test_optimize_applies_full_category_set(self): + mgr = self._mgr() + mgr.register_categorical_tag("weather", ["rainy", "sunny", "cloudy", "snow"]) + df = pd.DataFrame( + {"sample_id": [1, 2], "origin": "train", "tag:weather": ["rainy", "sunny"]} + ).set_index("sample_id") + mgr.upsert_df(df, origin="train") + view = mgr.get_df_view() + col = view["tag:weather"] + self.assertTrue(isinstance(col.dtype, pd.CategoricalDtype)) + # Full registered set present even though only 2 values appear in data + self.assertEqual(set(col.dtype.categories), {"rainy", "sunny", "cloudy", "snow"}) + + +class TestCategoricalTagH5RoundTrip(unittest.TestCase): + def _store(self): + d = tempfile.mkdtemp() + return H5DataFrameStore(Path(d) / "data.h5") + + def test_registry_save_load(self): + store = self._store() + reg = {"weather": ["rainy", "sunny", "cloudy"], "quality": ["high", "low"]} + store.save_tag_registry(reg) + loaded = store.load_tag_registry() + self.assertEqual(loaded, reg) + + def test_round_trip_preserves_unused_categories(self): + store = self._store() + store.save_tag_registry({"weather": ["rainy", "sunny", "cloudy", "snow"]}) + + idx = pd.MultiIndex.from_arrays( + [["a", "b", "c"], [0, 0, 0]], names=["sample_id", "annotation_id"] + ) + df = pd.DataFrame( + { + "origin": ["train"] * 3, + "tag:weather": ["rainy", "sunny", "rainy"], # only 2 of 4 used + "tag:is_urban": [True, False, True], # boolean tag + }, + index=idx, + ) + store.upsert("train", df) + back = store.load("train") + + weather = back["tag:weather"] + self.assertTrue(isinstance(weather.dtype, pd.CategoricalDtype)) + self.assertEqual(set(weather.dtype.categories), {"rainy", "sunny", "cloudy", "snow"}) + values = back.set_index("sample_id")["tag:weather"].astype(str).to_dict() + self.assertEqual(values, {"a": "rainy", "b": "sunny", "c": "rainy"}) + # Boolean tag survives independently + self.assertIn("tag:is_urban", back.columns) + + def test_clear_value_becomes_unset(self): + store = self._store() + store.save_tag_registry({"weather": ["rainy", "sunny"]}) + idx = pd.MultiIndex.from_arrays([["a", "b"], [0, 0]], names=["sample_id", "annotation_id"]) + df = pd.DataFrame( + {"origin": ["train"] * 2, "tag:weather": ["rainy", None]}, index=idx + ) + store.upsert("train", df) + back = store.load("train").set_index(["sample_id", "annotation_id"]) + # 'b' had no value → unset (NaN), 'a' keeps its category + self.assertEqual(str(back.loc[("a", 0), "tag:weather"]), "rainy") + self.assertTrue(pd.isna(back.loc[("b", 0), "tag:weather"])) + + +class TestBooleanTagsStillWork(unittest.TestCase): + """Regression guards: the legacy boolean-tag path must be unaffected.""" + + def _store(self): + return H5DataFrameStore(Path(tempfile.mkdtemp()) / "data.h5") + + def test_boolean_tag_not_registered_categorical(self): + mgr = LedgeredDataFrameManager(enable_flushing_threads=False, enable_h5_persistence=False) + df = pd.DataFrame( + {"sample_id": [1, 2, 3], "origin": "train", "tag:is_urban": [True, False, True]} + ).set_index("sample_id") + mgr.upsert_df(df, origin="train") + self.assertEqual(mgr.get_categorical_tags(), {}) + + def test_boolean_tag_and_discarded_survive_reupsert(self): + # Regression: re-upserting must not stringify True/False into "True"/"False". + store = self._store() + idx = pd.MultiIndex.from_arrays([["a", "b"], [0, 0]], names=["sample_id", "annotation_id"]) + df = pd.DataFrame( + {"origin": ["train"] * 2, "tag:flagged": [True, False], "discarded": [False, True]}, + index=idx, + ) + store.upsert("train", df) + store.upsert("train", df) # merge path + back = store.load("train").set_index(["sample_id", "annotation_id"]) + self.assertTrue(bool(back.loc[("a", 0), "tag:flagged"])) + self.assertFalse(bool(back.loc[("b", 0), "tag:flagged"])) + self.assertFalse(bool(back.loc[("a", 0), "discarded"])) + self.assertTrue(bool(back.loc[("b", 0), "discarded"])) + + def test_boolean_and_categorical_coexist(self): + store = self._store() + store.save_tag_registry({"weather": ["rainy", "sunny"]}) + idx = pd.MultiIndex.from_arrays([["a", "b"], [0, 0]], names=["sample_id", "annotation_id"]) + df = pd.DataFrame( + {"origin": ["train"] * 2, "tag:weather": ["rainy", "sunny"], "tag:flagged": [True, False]}, + index=idx, + ) + store.upsert("train", df) + store.upsert("train", df) + back = store.load("train").set_index(["sample_id", "annotation_id"]) + self.assertTrue(isinstance(back["tag:weather"].dtype, pd.CategoricalDtype)) + self.assertEqual(str(back.loc[("a", 0), "tag:weather"]), "rainy") + self.assertTrue(bool(back.loc[("a", 0), "tag:flagged"])) + self.assertFalse(bool(back.loc[("b", 0), "tag:flagged"])) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/data/test_data_samples_with_ops.py b/tests/data/test_data_samples_with_ops.py new file mode 100644 index 00000000..702cd1f8 --- /dev/null +++ b/tests/data/test_data_samples_with_ops.py @@ -0,0 +1,783 @@ +""" +Unit tests for data_samples_with_ops.py + +Tests the DataSampleTrackingWrapper class, which wraps PyTorch datasets +and provides per-sample statistics tracking and tag-based labeling. +""" + +import os +import tempfile +import unittest +import numpy as np +import pandas as pd +import torch + +import weightslab as wl + +from torch.utils.data import Dataset +from unittest.mock import patch + +import weightslab.data.data_samples_with_ops as _dswo +from weightslab.data.data_samples_with_ops import ( + DataSampleTrackingWrapper, + _has_regex_symbols, + _match_column_patterns, + _filter_columns_with_patterns, +) + + +def _reset_global_uid_state(): + """Reset module-level UID globals to avoid cross-test contamination.""" + _dswo._UID_CNT = 0 + _dswo._GLOBAL_UID_REGISTRY.clear() + + +class SimpleDataset(Dataset): + """Simple dataset for testing.""" + + def __init__(self, size=10): + self.size = size + self.__name__ = "simple_dataset" + + def __len__(self): + return self.size + + def __getitem__(self, idx): + # Return random data with shape (3, 32, 32) to simulate images + data = np.random.randn(3, 32, 32).astype(np.float32) + uid = str(idx) # Consistent with string UID preference + label = idx % 10 # Simulate 10 classes + return data, uid, label + + +class TestHelperFunctions(unittest.TestCase): + """Test helper utility functions.""" + + def test_has_regex_symbols(self): + """Test regex symbol detection.""" + # True cases + self.assertTrue(_has_regex_symbols(".*")) + self.assertTrue(_has_regex_symbols("test.*")) + self.assertTrue(_has_regex_symbols("[abc]")) + self.assertTrue(_has_regex_symbols("(test)")) + self.assertTrue(_has_regex_symbols("test+")) + + # False cases + self.assertFalse(_has_regex_symbols("test")) + self.assertFalse(_has_regex_symbols("test_column")) + self.assertFalse(_has_regex_symbols("123")) + + def test_match_column_patterns(self): + """Test column pattern matching.""" + # Exact match + self.assertTrue(_match_column_patterns("test_col", ["test_col"])) + self.assertTrue(_match_column_patterns("exact", ["exact", "other"])) + + # Regex match + self.assertTrue(_match_column_patterns("test_1", ["test_.*"])) + self.assertTrue(_match_column_patterns("feature_loss", [".*_loss"])) + + # No match + self.assertFalse(_match_column_patterns("column", ["other"])) + self.assertFalse(_match_column_patterns("test", [".*_loss"])) + + def test_filter_columns_with_patterns(self): + """Test column filtering by patterns.""" + columns = ["loss", "loss_train", "accuracy", "test_accuracy", "feature_map"] + + # Exact patterns + result = _filter_columns_with_patterns(columns, ["loss"]) + self.assertEqual(result, ["loss"]) + + # Regex patterns - note: the regex is correctly anchored, so ".*accuracy" matches "accuracy" and "test_accuracy" + result = _filter_columns_with_patterns(columns, [".*accuracy"]) + self.assertIn("accuracy", result) + self.assertIn("test_accuracy", result) + + # Multiple patterns - "loss" and "accuracy" should match "loss", "accuracy", "test_accuracy" = 2 matches (not 3 - loss_train is separate) + result = _filter_columns_with_patterns(columns, ["loss", "accuracy"]) + # "loss" matches exactly "loss" (1), "accuracy" matches "accuracy" and "test_accuracy" (2) = 2 total + self.assertGreaterEqual(len(result), 2) + + # Empty result + result = _filter_columns_with_patterns(columns, ["nonexistent"]) + self.assertEqual(result, []) + + +class TestDataSampleTrackingWrapperInit(unittest.TestCase): + """Test initialization and basic properties.""" + + def setUp(self): + """Create a temporary directory for logs.""" + _reset_global_uid_state() + self.temp_dir = tempfile.mkdtemp() + self.dataset = SimpleDataset(size=10) + + # Initialize HP + parameters = { + 'flush_interval': 3.0, + 'flush_max_rows': 100, + 'enable_h5': True, + 'enable_flush': True + } + wl.watch_or_edit( + parameters, + flag="hyperparameters", + defaults=parameters, + poll_interval=1.0, + ) + + def tearDown(self): + """Clean up temporary files.""" + import shutil + from weightslab.backend.ledgers import clear_all + clear_all() + if os.path.exists(self.temp_dir): + shutil.rmtree(self.temp_dir) + + def test_initialization_with_valid_params(self): + """Test wrapper initialization with valid parameters.""" + + wrapper = DataSampleTrackingWrapper( + wrapped_dataset=self.dataset, + root_log_dir=self.temp_dir, + is_training=True, + loader_name="train", + enable_h5_persistence=False, + compute_hash=False, + ) + + self.assertEqual(len(wrapper), len(self.dataset)) + self.assertEqual(wrapper.loader_name, "train") + self.assertTrue(wrapper.is_training) + self.assertIsNotNone(wrapper.unique_ids) + + def test_length_matches_dataset(self): + """Test that wrapper length matches dataset length.""" + + dataset_size = 20 + dataset = SimpleDataset(size=dataset_size) + wrapper = DataSampleTrackingWrapper( + wrapped_dataset=dataset, + root_log_dir=self.temp_dir, + enable_h5_persistence=False, + compute_hash=False, + ) + + self.assertEqual(len(wrapper), dataset_size) + + +class TestDataSampleTrackingWrapperGetItem(unittest.TestCase): + """Test __getitem__ and data retrieval.""" + + def setUp(self): + """Create a temporary directory for logs.""" + _reset_global_uid_state() + self.temp_dir = tempfile.mkdtemp() + self.dataset = SimpleDataset(size=10) + + # Initialize HP + parameters = { + 'flush_interval': 3.0, + 'flush_max_rows': 100, + 'enable_h5': True, + 'enable_flush': True + } + wl.watch_or_edit( + parameters, + flag="hyperparameters", + defaults=parameters, + poll_interval=1.0, + ) + + def tearDown(self): + """Clean up temporary files.""" + import shutil + from weightslab.backend.ledgers import clear_all + clear_all() + if os.path.exists(self.temp_dir): + shutil.rmtree(self.temp_dir) + + def test_getitem_returns_data_and_id(self): + """Test that __getitem__ returns (data, id, label, ...).""" + + wrapper = DataSampleTrackingWrapper( + wrapped_dataset=self.dataset, + root_log_dir=self.temp_dir, + enable_h5_persistence=False, + compute_hash=False, + use_tags=False, + ) + + # Get first item + result = wrapper[0] + + # Should return tuple with (data, id, target, ...) + self.assertIsInstance(result, tuple) + self.assertGreaterEqual(len(result), 3) # data, id, target at minimum + + # First element should be numpy array or tensor + self.assertTrue(isinstance(result[0], (np.ndarray, torch.Tensor))) + + # Second element should be a numeric UID + self.assertTrue(isinstance(result[1], (str, int))) + + +class TestDataSampleTrackingWrapperTagBasedLabeling(unittest.TestCase): + """Test tag-based labeling functionality.""" + + def setUp(self): + """Create a temporary directory for logs.""" + _reset_global_uid_state() + self.temp_dir = tempfile.mkdtemp() + self.dataset = SimpleDataset(size=10) + + # Initialize HP + parameters = { + 'flush_interval': 3.0, + 'flush_max_rows': 100, + 'enable_h5': True, + 'enable_flush': True + } + wl.watch_or_edit( + parameters, + flag="hyperparameters", + defaults=parameters, + poll_interval=1.0, + ) + + def tearDown(self): + """Clean up temporary files.""" + import shutil + from weightslab.backend.ledgers import clear_all + clear_all() + if os.path.exists(self.temp_dir): + shutil.rmtree(self.temp_dir) + + def test_binary_tag_labeling(self): + """Test tag-based labeling with individual boolean tag columns. + + Tags are now stored as individual boolean columns (tags_tagname) instead of + a single comma-separated string. This test verifies that tagging creates the + appropriate columns and that the tags are correctly stored and retrieved. + """ + + self.temp_dir = tempfile.mkdtemp() + tags_mapping = {"target_tag": 1, "non_target_tag": 0} + wrapper = DataSampleTrackingWrapper( + wrapped_dataset=self.dataset, + root_log_dir=self.temp_dir, + enable_h5_persistence=False, + compute_hash=False, + use_tags=True, + tags_mapping=tags_mapping, + ) + + # Get sample IDs + sample_id_0 = wrapper.unique_ids[0] + sample_id_1 = wrapper.unique_ids[1] + + # Verify no tag columns exist initially + df = wrapper.get_dataframe() + tag_columns_before = [col for col in df.columns if col.startswith("tag:")] + self.assertEqual(len(tag_columns_before), 0, "No tag columns should exist initially") + + # Set target_tag for first sample + wrapper.set(sample_id=sample_id_0, stat_name="tags", value='target_tag') + + # Verify tag column was created + df = wrapper.get_dataframe() + self.assertIn('tag:target_tag', df.columns, "tag:target_tag column should exist") + # With multi-index, use .xs() or get first value if Series returned + tag_value_0 = df.loc[sample_id_0, 'tag:target_tag'] + if isinstance(tag_value_0, pd.Series): + tag_value_0 = tag_value_0.iloc[0] + self.assertEqual(tag_value_0, 1, "target_tag should be set to 1") + + # Set non_target_tag for second sample + wrapper.set(sample_id=sample_id_1, stat_name="tags", value='non_target_tag') + + # Verify both tag columns exist + df = wrapper.get_dataframe() + self.assertIn('tag:target_tag', df.columns) + self.assertIn('tag:non_target_tag', df.columns) + tag_value_1 = df.loc[sample_id_1, 'tag:non_target_tag'] + if isinstance(tag_value_1, pd.Series): + tag_value_1 = tag_value_1.iloc[0] + self.assertEqual(tag_value_1, 1) + + def test_binary_tag_labeling_single_tag(self): + """Test binary tag-based labeling with a single target tag. + + When tags_mapping has only 1 tag, it's binary classification: + - tag matches → 1 + - tag doesn't match (or no tags) → 0 + """ + + self.temp_dir = tempfile.mkdtemp() + tags_mapping = {"target_tag": 1} # Binary: only 1 tag in mapping + wrapper = DataSampleTrackingWrapper( + wrapped_dataset=self.dataset, + root_log_dir=self.temp_dir, + enable_h5_persistence=False, + compute_hash=False, + use_tags=True, + tags_mapping=tags_mapping, + ) + + sample_id_0 = wrapper.unique_ids[0] + + # Set target_tag on sample 0 + wrapper.set(sample_id=sample_id_0, stat_name="tags", value='target_tag') + + # Verify tags were set + df = wrapper.get_dataframe() + self.assertIn('tag:target_tag', df.columns) + # Handle multi-index - loc returns Series + tag_value = df.loc[sample_id_0, 'tag:target_tag'] + if isinstance(tag_value, pd.Series): + tag_value = tag_value.iloc[0] + self.assertEqual(tag_value, 1) + + def test_tag_parsing_comma_and_semicolon(self): + """Test that tags can be separated by commas, semicolons, or both. + + The tag parsing should handle: + - "tag1,tag2,tag3" → creates tag:tag1, tag:tag2, tag:tag3 + - "tag1;tag2;tag3" → creates tag:tag1, tag:tag2, tag:tag3 + - "tag1, tag2; tag3" → creates tag:tag1, tag:tag2, tag:tag3 (trims whitespace) + """ + + self.temp_dir = tempfile.mkdtemp() + tags_mapping = {"tag1": 1, "tag2": 2, "tag3": 3} + wrapper = DataSampleTrackingWrapper( + wrapped_dataset=self.dataset, + root_log_dir=self.temp_dir, + enable_h5_persistence=False, + compute_hash=False, + use_tags=True, + tags_mapping=tags_mapping, + ) + + sample_id = wrapper.unique_ids[0] + + # Test comma-separated tags + wrapper.set(sample_id=sample_id, stat_name="tags", value='tag1,tag2,tag3') + df = wrapper.get_dataframe() + for tag in ['tag:tag1', 'tag:tag2', 'tag:tag3']: + val = df.loc[sample_id, tag] + if isinstance(val, pd.Series): + val = val.iloc[0] + self.assertEqual(val, 1) + + # Test semicolon-separated tags on different sample + sample_id_2 = wrapper.unique_ids[2] + wrapper.set(sample_id=sample_id_2, stat_name="tags", value='tag1;tag2;tag3') + df = wrapper.get_dataframe() + for tag in ['tag:tag1', 'tag:tag2', 'tag:tag3']: + val = df.loc[sample_id_2, tag] + if isinstance(val, pd.Series): + val = val.iloc[0] + self.assertEqual(val, 1) + +class TestDataSampleTrackingWrapperDenylist(unittest.TestCase): + """Test denylisting and allowlisting functionality.""" + + def setUp(self): + """Create a temporary directory for logs.""" + _reset_global_uid_state() + self.temp_dir = tempfile.mkdtemp() + self.dataset = SimpleDataset(size=10) + + # Initialize HP + parameters = { + 'flush_interval': 3.0, + 'flush_max_rows': 100, + 'enable_h5': True, + 'enable_flush': True + } + wl.watch_or_edit( + parameters, + flag="hyperparameters", + defaults=parameters, + poll_interval=1.0, + ) + + def tearDown(self): + """Clean up temporary files.""" + import shutil + from weightslab.backend.ledgers import clear_all + clear_all() + if os.path.exists(self.temp_dir): + shutil.rmtree(self.temp_dir) + + def test_denylist_samples(self): + """Test denylisting samples.""" + + wrapper = DataSampleTrackingWrapper( + wrapped_dataset=self.dataset, + root_log_dir=self.temp_dir, + enable_h5_persistence=False, + compute_hash=False, + ) + + # Get first few sample IDs + denied_ids = set(wrapper.unique_ids[:3]) + + wrapper.denylist_samples(denied_ids) + + # Check denied count was updated + self.assertEqual(wrapper.denied_sample_cnt, len(denied_ids)) + + def test_allowlist_samples(self): + """Test allowlisting samples.""" + + wrapper = DataSampleTrackingWrapper( + wrapped_dataset=self.dataset, + root_log_dir=self.temp_dir, + enable_h5_persistence=False, + compute_hash=False, + ) + + # First deny some samples + denied_ids = set(wrapper.unique_ids[:3]) + wrapper.denylist_samples(denied_ids) + self.assertEqual(wrapper.denied_sample_cnt, len(denied_ids)) + + # Then allow them back + wrapper.allowlist_samples(denied_ids) + + # If allow was successful, denied_sample_cnt should be updated + # (depends on mock behavior of get_df_view) + + def test_denylist_clear(self): + """Test clearing all denylists.""" + + wrapper = DataSampleTrackingWrapper( + wrapped_dataset=self.dataset, + root_log_dir=self.temp_dir, + enable_h5_persistence=False, + compute_hash=False, + ) + + # Deny all samples + all_ids = set(wrapper.unique_ids) + wrapper.denylist_samples(all_ids) + self.assertEqual(wrapper.denied_sample_cnt, len(all_ids)) + + # Clear denials by passing None + wrapper.denylist_samples(None) + self.assertEqual(wrapper.denied_sample_cnt, 0) + + +class TestDataSampleTrackingWrapperStateDict(unittest.TestCase): + """Test state_dict and load_state_dict functionality.""" + + def setUp(self): + """Create a temporary directory for logs.""" + _reset_global_uid_state() + self.temp_dir = tempfile.mkdtemp() + self.dataset = SimpleDataset(size=5) + + # Initialize HP + parameters = { + 'flush_interval': 3.0, + 'flush_max_rows': 100, + 'enable_h5': True, + 'enable_flush': True + } + wl.watch_or_edit( + parameters, + flag="hyperparameters", + defaults=parameters, + poll_interval=1.0, + ) + + def tearDown(self): + """Clean up temporary files.""" + import shutil + from weightslab.backend.ledgers import clear_all + clear_all() + if os.path.exists(self.temp_dir): + shutil.rmtree(self.temp_dir) + + def test_state_dict_structure(self): + """Test that state_dict has correct structure.""" + + wrapper = DataSampleTrackingWrapper( + wrapped_dataset=self.dataset, + root_log_dir=self.temp_dir, + enable_h5_persistence=False, + compute_hash=False, + ) + + state = wrapper.state_dict() + + # Check structure + self.assertIn("blockd_samples", state) + self.assertIn("sample_statistics", state) + self.assertIsInstance(state["blockd_samples"], int) + self.assertIsInstance(state["sample_statistics"], dict) + + +class TestDataSampleTrackingWrapperUtilities(unittest.TestCase): + """Test utility methods.""" + + def setUp(self): + """Create a temporary directory for logs.""" + _reset_global_uid_state() + self.temp_dir = tempfile.mkdtemp() + self.dataset = SimpleDataset(size=10) + + # Initialize HP + parameters = { + 'flush_interval': 3.0, + 'flush_max_rows': 100, + 'enable_h5': True, + 'enable_flush': True + } + wl.watch_or_edit( + parameters, + flag="hyperparameters", + defaults=parameters, + poll_interval=1.0, + ) + + def tearDown(self): + """Clean up temporary files.""" + import shutil + from weightslab.backend.ledgers import clear_all + clear_all() + if os.path.exists(self.temp_dir): + shutil.rmtree(self.temp_dir) + + def test_get_sample_id_at_index(self): + """Test retrieving sample ID at given index.""" + + wrapper = DataSampleTrackingWrapper( + wrapped_dataset=self.dataset, + root_log_dir=self.temp_dir, + enable_h5_persistence=False, + compute_hash=False, + ) + + # Get sample ID at index 0 + sample_id = wrapper.get_sample_id_at_index(0) + self.assertEqual(sample_id, wrapper.unique_ids[0]) + + def test_get_index_from_sample_id(self): + """Test retrieving index from sample ID.""" + + wrapper = DataSampleTrackingWrapper( + wrapped_dataset=self.dataset, + root_log_dir=self.temp_dir, + enable_h5_persistence=False, + compute_hash=False, + ) + + # Get index from first sample ID + sample_id = wrapper.unique_ids[0] + index = wrapper.get_index_from_sample_id(sample_id) + self.assertEqual(index, 0) + + def test_infer_num_classes_from_dataset(self): + """Test inferring number of classes from wrapped dataset.""" + + # Create a dataset with num_classes attribute + dataset = SimpleDataset(size=10) + dataset.num_classes = 10 + + wrapper = DataSampleTrackingWrapper( + wrapped_dataset=dataset, + root_log_dir=self.temp_dir, + enable_h5_persistence=False, + compute_hash=False, + ) + + num_classes = wrapper.infer_num_classes() + self.assertEqual(num_classes, 10) + + def test_infer_num_classes_binary_tags(self): + """Test inferring num_classes with binary tag mapping.""" + + wrapper = DataSampleTrackingWrapper( + wrapped_dataset=self.dataset, + root_log_dir=self.temp_dir, + enable_h5_persistence=False, + compute_hash=False, + use_tags=True, + tags_mapping={"target": 1}, + ) + + num_classes = wrapper.infer_num_classes() + self.assertEqual(num_classes, 2) + + +class TestDataSampleTrackingWrapperDuplicateDetection(unittest.TestCase): + """Test duplicate sample detection and removal.""" + + def setUp(self): + """Create a temporary directory for logs.""" + _reset_global_uid_state() + self.temp_dir = tempfile.mkdtemp() + + # Initialize HP + parameters = { + 'flush_interval': 3.0, + 'flush_max_rows': 100, + 'enable_h5': True, + 'enable_flush': True + } + wl.watch_or_edit( + parameters, + flag="hyperparameters", + defaults=parameters, + poll_interval=1.0, + ) + + def tearDown(self): + """Clean up temporary files.""" + import shutil + from weightslab.backend.ledgers import clear_all + clear_all() + if os.path.exists(self.temp_dir): + shutil.rmtree(self.temp_dir) + + @patch('weightslab.data.data_samples_with_ops.array_id_2bytes') + def test_duplicate_detection_with_hash(self, mock_hash): + """Test that duplicate samples are detected and removed.""" + + # Create dataset with duplicates + dataset = SimpleDataset(size=5) + + # Mock hash function to create duplicates + # Return same hash for first two samples + hash_values = [100, 100, 101, 102, 103] + mock_hash.side_effect = hash_values + + wrapper = DataSampleTrackingWrapper( + wrapped_dataset=dataset, + root_log_dir=self.temp_dir, + enable_h5_persistence=False, + compute_hash=True, + ) + + # Should have removed one duplicate + # The wrapper should now have 4 unique samples instead of 5 + # (though actual behavior depends on how Subset is used) + self.assertLessEqual(len(wrapper.unique_ids), len(dataset)) + + +class TestDataSampleTrackingWrapperEquality(unittest.TestCase): + """Test equality comparison.""" + + def setUp(self): + """Create a temporary directory for logs.""" + _reset_global_uid_state() + self.temp_dir = tempfile.mkdtemp() + self.dataset = SimpleDataset(size=10) + + # Initialize HP + parameters = { + 'flush_interval': 3.0, + 'flush_max_rows': 100, + 'enable_h5': True, + 'enable_flush': True + } + wl.watch_or_edit( + parameters, + flag="hyperparameters", + defaults=parameters, + poll_interval=1.0, + ) + + def tearDown(self): + """Clean up temporary files.""" + import shutil + if os.path.exists(self.temp_dir): + shutil.rmtree(self.temp_dir) + + def test_equality_same_wrapper(self): + """Test equality comparison of wrappers.""" + + wrapper1 = DataSampleTrackingWrapper( + wrapped_dataset=self.dataset, + root_log_dir=self.temp_dir, + enable_h5_persistence=False, + compute_hash=False, + ) + + wrapper2 = DataSampleTrackingWrapper( + wrapped_dataset=self.dataset, + root_log_dir=self.temp_dir, + enable_h5_persistence=False, + compute_hash=False, + ) + + # Both have same wrapped_dataset and same denied_count + self.assertTrue(wrapper1 == wrapper2) + + def test_equality_different_types(self): + """Test equality comparison with different types.""" + + wrapper = DataSampleTrackingWrapper( + wrapped_dataset=self.dataset, + root_log_dir=self.temp_dir, + enable_h5_persistence=False, + compute_hash=False, + ) + + # Compare with non-wrapper object + self.assertFalse(wrapper == "not a wrapper") + self.assertFalse(wrapper == 123) + + +class TestDataSampleTrackingWrapperAsRecords(unittest.TestCase): + """Test as_records functionality.""" + + def setUp(self): + """Create a temporary directory for logs.""" + _reset_global_uid_state() + self.temp_dir = tempfile.mkdtemp() + self.dataset = SimpleDataset(size=5) + + # Initialize HP + parameters = { + 'flush_interval': 3.0, + 'flush_max_rows': 100, + 'enable_h5': True, + 'enable_flush': True + } + wl.watch_or_edit( + parameters, + flag="hyperparameters", + defaults=parameters, + poll_interval=1.0, + ) + + def tearDown(self): + """Clean up temporary files.""" + import shutil + if os.path.exists(self.temp_dir): + shutil.rmtree(self.temp_dir) + + def test_as_records(self): + """Test converting DataFrame to records.""" + + wrapper = DataSampleTrackingWrapper( + wrapped_dataset=self.dataset, + root_log_dir=self.temp_dir, + enable_h5_persistence=False, + compute_hash=False, + ) + + records = wrapper.as_records() + + self.assertIsInstance(records, list) + self.assertGreater(len(records), 0) + for record in records: + self.assertIsInstance(record, dict) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/data/test_data_service_metadata_copy.py b/tests/data/test_data_service_metadata_copy.py new file mode 100644 index 00000000..c8185d6e --- /dev/null +++ b/tests/data/test_data_service_metadata_copy.py @@ -0,0 +1,97 @@ +import unittest + +import pandas as pd + +from weightslab.trainer.services.data_service import ( + normalize_metadata_copy_source_name, + build_metadata_copy_column_names, + duplicate_metadata_column_in_dataframe, + is_copy_metadata_column_name, + is_protected_metadata_name, +) + + +class TestDataServiceMetadataCopyHelpers(unittest.TestCase): + def test_normalize_source_name_strips_hash_prefix(self): + normalized = normalize_metadata_copy_source_name("oldhash@my metric", "newhash") + self.assertEqual(normalized, "my_metric") + + def test_normalize_source_name_removes_trailing_numeric_suffix(self): + normalized = normalize_metadata_copy_source_name("newhash@quality_score_7", "newhash") + self.assertEqual(normalized, "quality_score") + + def test_normalize_source_name_keeps_non_numeric_suffix(self): + normalized = normalize_metadata_copy_source_name("newhash@quality_score_final", "newhash") + self.assertEqual(normalized, "quality_score_final") + + def test_build_names_starts_with_index_one(self): + existing_columns = ["sample_id", "origin", "something_else"] + backend_name = build_metadata_copy_column_names( + existing_columns, + "abc123", + "oldhash@quality_score", + ) + self.assertEqual(backend_name, "quality_score_1@abc123") + + def test_build_names_increments_index_when_existing(self): + existing_columns = [ + "quality_score_1@abc123", + "quality_score_2@abc123", + "other_metric_1@abc123", + ] + backend_name = build_metadata_copy_column_names( + existing_columns, + "abc123", + "quality_score", + ) + self.assertEqual(backend_name, "quality_score_3@abc123") + + def test_duplicate_column_copies_values(self): + df = pd.DataFrame( + { + "origin": ["train", "train", "val"], + "oldhash@quality_score": [0.1, 0.2, 0.9], + }, + index=["1", "2", "3"], + ) + + duplicated, backend_name = duplicate_metadata_column_in_dataframe( + df, + source_column="oldhash@quality_score", + experiment_hash="newhash", + ) + + self.assertEqual(backend_name, "quality_score_1@newhash") + self.assertIn(backend_name, duplicated.columns) + self.assertListEqual( + duplicated[backend_name].tolist(), + duplicated["oldhash@quality_score"].tolist(), + ) + + def test_duplicate_raises_for_missing_source(self): + df = pd.DataFrame({"origin": ["train"]}, index=["1"]) + + with self.assertRaises(KeyError): + duplicate_metadata_column_in_dataframe( + df, + source_column="missing@metadata", + experiment_hash="abc123", + ) + + def test_remove_action_accepts_only_copy_metadata_columns(self): + self.assertTrue(is_copy_metadata_column_name("quality_score_1@abc123")) + self.assertTrue(is_copy_metadata_column_name("foo_bar_99@exp")) + self.assertFalse(is_copy_metadata_column_name("quality_score@abc123")) + self.assertFalse(is_copy_metadata_column_name("origin")) + + def test_remove_action_rejects_protected_metadata_columns(self): + self.assertTrue(is_protected_metadata_name("sample_id")) + self.assertTrue(is_protected_metadata_name("origin")) + self.assertTrue(is_protected_metadata_name("tag:hard_example")) + self.assertTrue(is_protected_metadata_name("signal:confidence")) + self.assertTrue(is_protected_metadata_name("signals//loss")) + self.assertFalse(is_protected_metadata_name("quality_score_1@abc123")) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/data/test_data_utils_unit.py b/tests/data/test_data_utils_unit.py new file mode 100644 index 00000000..a51888d9 --- /dev/null +++ b/tests/data/test_data_utils_unit.py @@ -0,0 +1,117 @@ +import unittest + +import numpy as np + +from weightslab.data import data_utils as du + + +class _DatasetWrapper: + def __init__(self, rows): + self.rows = rows + + def __getitem__(self, idx): + return self.rows[idx] + + +class _DatasetWithIndex: + def __init__(self, rows): + self.wrapped_dataset = _DatasetWrapper(rows) + + def get_index_from_sample_id(self, sample_id): + return int(sample_id) + + +class _SplitObj: + pass + + +class TestDataUtilsUnit(unittest.TestCase): + def test_pattern_matching_and_cache(self): + du._PATTERN_CACHE.clear() + self.assertIsNone(du._get_compiled_pattern("[")) + self.assertIn("[", du._PATTERN_CACHE) + + cols = ["loss", "signals//train_loss", "acc"] + out = du._filter_columns_by_patterns(cols, ["loss", ".*train_loss$"]) + self.assertEqual(out, ["loss", "signals//train_loss"]) + self.assertTrue(du._matches_pattern("signals//train_loss", [".*train_loss$"])) + self.assertFalse(du._matches_pattern("metric", ["^loss$"])) + + def test_split_detection_and_downsample(self): + ds = _SplitObj() + ds.train = True + self.assertEqual(du._detect_dataset_split(ds), "train") + + ds2 = _SplitObj() + ds2.train = False + ds2.split = " Val " + self.assertEqual(du._detect_dataset_split(ds2), "val") + + ds3 = _SplitObj() + ds3.mode = "TEST" + self.assertEqual(du._detect_dataset_split(ds3), "test") + + arr2d = np.arange(10000).reshape(100, 100) + self.assertLessEqual(max(du._downsample_nn(arr2d, max_hw=20).shape), 20) + + arr3d_chw = np.zeros((3, 100, 80), dtype=np.float32) + out_chw = du._downsample_nn(arr3d_chw, max_hw=20) + self.assertEqual(out_chw.shape[0], 3) + + def test_to_numpy_and_mask_helpers(self): + self.assertEqual(du.to_numpy_safe(3).tolist(), [3]) + self.assertEqual(du.to_numpy_safe([1, 2]).tolist(), [1, 2]) + + bboxes = np.array([[1, 1, 3, 3, 2]], dtype=np.float32) + raw_data = (np.zeros((5, 5, 3), dtype=np.uint8),) + mask = du.get_mask(bboxes, raw_data=raw_data) + self.assertEqual(mask.shape, (5, 5)) + self.assertEqual(int(mask[1, 1]), 2) + + def test_label_metadata_uid_and_volumetric_helpers(self): + rows = [( + np.zeros((4, 4, 3), dtype=np.uint8), + "uid-0", + np.array([[0, 0, 2, 2]], dtype=np.int64), + {"classes": np.array([5], dtype=np.int64), "source": "a"}, + )] + dataset = _DatasetWithIndex(rows) + + label = du.load_label(dataset, "0") + self.assertEqual(label.shape, (1, 5)) + self.assertEqual(int(label[0, 4]), 5) + + metadata_dataset = _DatasetWithIndex([( + np.zeros((4, 4, 3), dtype=np.uint8), + "uid-0", + np.array([1], dtype=np.int64), + {"source": "a"}, + {"fold": "train"}, + )]) + metadata = du.load_metadata(metadata_dataset, "0") + self.assertEqual(metadata["source"], "a") + self.assertEqual(metadata["fold"], "train") + + uid = du.load_uid(dataset, "0") + self.assertEqual(uid, "uid-0") + + vol = np.zeros((2, 8, 8, 1), dtype=np.float32) + sliced = du._extract_slice_from_4d(vol) + self.assertEqual(sliced.shape, (8, 8, 1)) + + def test_load_raw_image_array_and_invalid_channels(self): + rows = [(np.zeros((2, 6, 6, 1), dtype=np.float32),)] + dataset = _DatasetWithIndex(rows) + arr, is_vol, original_shape, middle = du.load_raw_image_array(dataset, 0) + self.assertTrue(is_vol) + self.assertEqual(tuple(original_shape), (2, 6, 6, 1)) + self.assertEqual(arr.shape, (2, 6, 6, 1)) + self.assertEqual(middle.mode, "L") + + bad = _DatasetWithIndex([(np.zeros((6, 6, 2), dtype=np.uint8),)]) + with self.assertRaises(ValueError): + du.load_raw_image(bad, 0) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/data/test_dataframe_data_invariants.py b/tests/data/test_dataframe_data_invariants.py new file mode 100644 index 00000000..428dde67 --- /dev/null +++ b/tests/data/test_dataframe_data_invariants.py @@ -0,0 +1,243 @@ +"""Regression tests pinning down core data invariants of LedgeredDataFrameManager. + +Each test guards a behavior that is easy to silently break: + +* bounding-box targets stay inline coordinates (never rasterized / spilled to H5), +* per-instance arrays (segmentation masks) round-trip distinctly through the array + store on resume, +* resuming an experiment restores the persisted instance rows onto a freshly + registered (sample-row-only) loader, +* a sample whose instances collapsed onto annotation_id 0 is repaired to 0..N, +* memory optimization downcasts signal columns to float32 and turns empty OBJECT + cells into None while keeping categorical columns categorical. +""" +import tempfile +import unittest +from pathlib import Path + +import numpy as np +import pandas as pd + +from weightslab.data.dataframe_manager import LedgeredDataFrameManager +from weightslab.data.h5_dataframe_store import H5DataFrameStore +from weightslab.data.h5_array_store import H5ArrayStore +from weightslab.data.array_proxy import ArrayH5Proxy +from weightslab.data.sample_stats import SampleStats + +TARGET = SampleStats.Ex.TARGET.value +ORIGIN = SampleStats.Ex.ORIGIN.value + + +def _write_legacy_checkpoint(store_path: Path, origin: str, df: pd.DataFrame): + """Write a genuine PRE-multi-index checkpoint frame to disk. + + Reproduces exactly what the old writer produced: a single-level ``sample_id`` + index (no ``annotation_id``), object columns stringified, under the store's + per-origin key (``/stats_``). Bypasses the current store.upsert so the + on-disk bytes are genuinely legacy — not silently promoted by the new writer. + """ + legacy = df.copy().set_index("sample_id") + legacy.columns = [str(c).replace('/', '__SLASH__') for c in legacy.columns] + for c in legacy.select_dtypes(include=['object']).columns: + legacy[c] = legacy[c].astype(str) + with pd.HDFStore(str(store_path), mode="a") as h: + h.put(f"/stats_{origin}", legacy, format="table", data_columns=True) + + +def _mgr(persist=False, tmp=None): + mgr = LedgeredDataFrameManager(enable_flushing_threads=False, enable_h5_persistence=persist) + if persist: + mgr.set_store(H5DataFrameStore(Path(tmp) / "data.h5")) + return mgr + + +class TestBoundingBoxStaysInline(unittest.TestCase): + def test_bbox_target_not_rasterized_or_proxied(self): + tmp = tempfile.mkdtemp() + mgr = _mgr(persist=True, tmp=tmp) + bbox = np.array([[1, 1, 3, 3, 2], [5, 5, 9, 9, 1]], dtype=np.float32) # (N, 5) + mgr.register_split("train", [{"sample_id": "d", "origin": "train", TARGET: bbox}]) + mgr.flush() + cell = mgr._df.loc[("d", 0), TARGET] + self.assertNotIsInstance(cell, ArrayH5Proxy) # not spilled to array H5 + np.testing.assert_array_equal(np.asarray(cell), bbox) # not rasterized + + def test_dense_mask_still_proxied(self): + tmp = tempfile.mkdtemp() + mgr = _mgr(persist=True, tmp=tmp) + mask = np.full((32, 32), 7, dtype=np.uint8) + mgr.register_split("train", [{"sample_id": "s", "origin": "train", TARGET: mask}]) + mgr.flush() + self.assertIsInstance(mgr._df.loc[("s", 0), TARGET], ArrayH5Proxy) + + +class TestPerInstanceArrayRoundTrip(unittest.TestCase): + def test_instance_masks_persist_and_load_distinctly(self): + tmp = tempfile.mkdtemp() + store_path = Path(tmp) / "data.h5" + masks = [np.full((32, 32), i, np.uint8) for i in (1, 2, 3)] + + mgr = LedgeredDataFrameManager(enable_flushing_threads=False, enable_h5_persistence=True) + mgr.set_store(H5DataFrameStore(store_path)) + mgr.register_split("train", [{"sample_id": "7", "origin": "train", TARGET: masks}]) + mgr.flush() + + # In-memory: each instance row holds a DISTINCT proxy path. + paths = {a: mgr._df.loc[("7", a), TARGET].path_ref for a in (1, 2, 3)} + self.assertEqual(len(set(paths.values())), 3) + + # Reload into a fresh manager (resume) and read each instance back. + mgr2 = LedgeredDataFrameManager(enable_flushing_threads=False, enable_h5_persistence=True) + mgr2.set_store(H5DataFrameStore(store_path)) + mgr2.register_split("train", [{"sample_id": "7", "origin": "train"}]) + v = mgr2.get_df_view() + for aid, expected in [(1, 1), (2, 2), (3, 3)]: + cell = v.loc[("7", aid), TARGET] + arr = cell.load() if isinstance(cell, ArrayH5Proxy) else np.asarray(cell) + self.assertEqual(arr.shape, (32, 32)) + self.assertEqual(int(np.median(arr)), expected) + + +class TestResumeRestoresInstanceRows(unittest.TestCase): + def test_instance_rows_and_signals_restored_on_resume(self): + tmp = tempfile.mkdtemp() + store_path = Path(tmp) / "data.h5" + + prev = LedgeredDataFrameManager(enable_flushing_threads=False, enable_h5_persistence=True) + prev.set_store(H5DataFrameStore(store_path)) + prev.register_split("train", [{"sample_id": "5", "origin": "train", + TARGET: [np.ones((2, 2)), np.full((2, 2), 2)]}]) + prev.enqueue_instance_batch(sample_ids=["5", "5"], annotation_ids=[1, 2], + losses={"signals//iou": np.array([0.7, 0.9])}, + step=1) + prev.flush() + + # Fresh run: register only the sample row (preload_labels=False style). + new = LedgeredDataFrameManager(enable_flushing_threads=False, enable_h5_persistence=True) + new.set_store(H5DataFrameStore(store_path)) + new.register_split("train", [{"sample_id": "5", "origin": "train"}]) + v = new.get_df_view() + + self.assertEqual(sorted(v.loc["5"].index.tolist()), [0, 1, 2]) + self.assertAlmostEqual(float(v.loc[("5", 1), "signals//iou"]), 0.7, places=5) + self.assertAlmostEqual(float(v.loc[("5", 2), "signals//iou"]), 0.9, places=5) + # Sample-level origin only on the sample row; instance rows stay clean. + self.assertEqual(v.loc[("5", 0), ORIGIN], "train") + + +class TestMultiInstanceRegistration(unittest.TestCase): + def test_list_target_expands_to_distinct_instance_rows(self): + """The supported multi-instance path: a single record whose target is a + LIST of array-likes expands to one sample row + N distinct instance rows.""" + mgr = LedgeredDataFrameManager(enable_flushing_threads=False, enable_h5_persistence=False) + mgr.register_split("train", [ + {"sample_id": "12", "origin": "train", + TARGET: [np.full((2, 2), i) for i in range(6)]}, # 6 instances + {"sample_id": "5", "origin": "train"}, # single-target sample → just the sample row + ]) + v = mgr.get_df_view() + # sample row (0) + 6 distinct instance rows (1..6) + self.assertEqual(sorted(v.loc["12"].index.tolist()), [0, 1, 2, 3, 4, 5, 6]) + self.assertEqual(sorted(v.loc["5"].index.tolist()), [0]) + self.assertFalse(v.index.has_duplicates) + + +class TestMemoryOptimizationInvariants(unittest.TestCase): + def test_signals_float32_and_object_nan_to_none_and_origin_categorical(self): + tmp = tempfile.mkdtemp() + mgr = LedgeredDataFrameManager(enable_flushing_threads=False, enable_h5_persistence=True) + mgr.set_store(H5DataFrameStore(Path(tmp) / "data.h5")) + masks = [np.full((4, 4), i, np.uint8) for i in (1, 2, 3)] + # Many samples across two origins so the compression ratio (n_unique / + # n_samples) is well under 0.5 and 'origin' is categoricalized. + mgr.register_split("train", [{"sample_id": str(i), "origin": "train"} for i in range(10)] + + [{"sample_id": "7", "origin": "train", TARGET: masks}]) + mgr.register_split("test", [{"sample_id": str(100 + i), "origin": "test"} for i in range(10)]) + mgr.enqueue_instance_batch(sample_ids=["7", "7", "7"], annotation_ids=[1, 2, 3], + losses={"signals//iou": np.array([0.1, 0.2, 0.3])}, + step=1) + mgr.flush() + df = mgr._df + + # signals downcast to float32 + self.assertEqual(df["signals//iou"].dtype, np.float32) + # origin categorical (memory-efficient); its instance-row missing is categorical NaN + import pandas as pd + self.assertIsInstance(df[ORIGIN].dtype, pd.CategoricalDtype) + # empty TARGET cell on the multi-instance sample row is None, not float nan + self.assertIsNone(df.loc[("7", 0), TARGET]) + + +class TestOldCheckpointCompatibility(unittest.TestCase): + """Old checkpoints (sandbox) were written with a SINGLE-LEVEL sample_id index + and NO annotation_id. Loading one with the current code must transparently + expand to the (sample_id, annotation_id=0) multi-index — instance_id generated + as 0 — with all sample data and arrays preserved.""" + + def test_load_old_single_level_dataframe_checkpoint(self): + tmp = tempfile.mkdtemp() + store_path = Path(tmp) / "data.h5" + + # 1) Genuine legacy on-disk frame: single-level sample_id index, no annotation_id. + old_df = pd.DataFrame({ + "sample_id": ["0", "1", "2"], + ORIGIN: ["train", "train", "train"], + "loss": [0.1, 0.2, 0.3], + SampleStats.Ex.DISCARDED.value: [False, True, False], + "tag:hard": [True, False, True], + }) + _write_legacy_checkpoint(store_path, "train", old_df) + + # Sanity: what's on disk really is single-level (no annotation_id column). + raw = H5DataFrameStore(store_path).load("train") + self.assertNotIn("annotation_id", raw.columns) + + # 2) Resume with the CURRENT store + manager (as the sandbox would on reload). + mgr = LedgeredDataFrameManager(enable_flushing_threads=False, enable_h5_persistence=True) + mgr.set_store(H5DataFrameStore(store_path)) + mgr.register_split("train", [{"sample_id": s, "origin": "train"} for s in ["0", "1", "2"]]) + v = mgr.get_df_view() + + # Expanded to (sample_id, annotation_id == 0) — instance_id generated as 0. + self.assertIsInstance(v.index, pd.MultiIndex) + self.assertEqual(list(v.index.names), ["sample_id", "annotation_id"]) + self.assertEqual(sorted(v.index.tolist()), [("0", 0), ("1", 0), ("2", 0)]) + # All legacy sample data preserved on the canonical rows. + self.assertAlmostEqual(float(v.loc[("0", 0), "loss"]), 0.1, places=5) + self.assertFalse(bool(v.loc[("0", 0), SampleStats.Ex.DISCARDED.value])) + self.assertTrue(bool(v.loc[("1", 0), SampleStats.Ex.DISCARDED.value])) + self.assertTrue(bool(v.loc[("0", 0), "tag:hard"])) + self.assertFalse(bool(v.loc[("1", 0), "tag:hard"])) + + def test_load_old_sample_level_array_via_bare_key(self): + """Old arrays.h5 stored each sample array at the bare '/sample_id/' + path (no composite annotation suffix). Those must still load via proxy.""" + tmp = tempfile.mkdtemp() + store_path = Path(tmp) / "data.h5" + + # Legacy array layout: '/0/prediction' (bare sample_id key). + arr_store = H5ArrayStore(Path(tmp) / "arrays.h5") + pred = np.full((32, 32), 9, np.uint8) + ref = arr_store.save_array("0", "prediction", pred, preserve_original=True) + self.assertTrue(ref.endswith(":/0/prediction")) # old-style bare key + + # Legacy dataframe references that array by its path-ref string. + _write_legacy_checkpoint(store_path, "train", pd.DataFrame({ + "sample_id": ["0"], + ORIGIN: ["train"], + SampleStats.Ex.PREDICTION.value: [ref], + })) + + mgr = LedgeredDataFrameManager(enable_flushing_threads=False, enable_h5_persistence=True) + mgr.set_store(H5DataFrameStore(store_path)) + mgr.register_split("train", [{"sample_id": "0", "origin": "train"}]) + v = mgr.get_df_view() + + cell = v.loc[("0", 0), SampleStats.Ex.PREDICTION.value] + arr = cell.load() if isinstance(cell, ArrayH5Proxy) else np.asarray(cell) + self.assertEqual(arr.shape, (32, 32)) + self.assertEqual(int(np.median(arr)), 9) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/data/test_dataframe_manager_unit.py b/tests/data/test_dataframe_manager_unit.py new file mode 100644 index 00000000..c12e91cc --- /dev/null +++ b/tests/data/test_dataframe_manager_unit.py @@ -0,0 +1,541 @@ +import hashlib +import tempfile +import unittest +from pathlib import Path +from unittest.mock import patch + +import numpy as np +import pandas as pd + +from weightslab.data.dataframe_manager import LedgeredDataFrameManager +from weightslab.data.h5_dataframe_store import H5DataFrameStore +from weightslab.data.sample_stats import SampleStats + + +class TestDataFrameManagerUnit(unittest.TestCase): + def test_sample_id_normalization_and_upsert(self): + mgr = LedgeredDataFrameManager(enable_flushing_threads=False, enable_h5_persistence=False) + self.assertEqual(mgr._normalize_sample_id(np.int64(7)), "7") + self.assertEqual(mgr._normalize_sample_id(b"abc"), "abc") + + df = pd.DataFrame([{"sample_id": 1, "origin": "train", "loss": 0.5}]).set_index("sample_id") + mgr.upsert_df(df, origin="train") + coerced = mgr._coerce_sample_id_for_index(1) + self.assertEqual(coerced, "1") + self.assertIn("1", mgr.get_df_view().index) + + def test_array_storage_and_safe_conversions(self): + mgr = LedgeredDataFrameManager(enable_flushing_threads=False, enable_h5_persistence=False) + + self.assertFalse(mgr._should_store_array_separately(np.array([1, 2, 3]))) + self.assertTrue(mgr._should_store_array_separately(np.zeros((30, 30), dtype=np.float32))) + + data = { + SampleStats.Ex.PREDICTION.value: np.zeros((30, 30), dtype=np.float32), + "other": "x", + } + arrays = mgr._extract_arrays_for_storage("1", data) + self.assertIn(SampleStats.Ex.PREDICTION.value, arrays) + + self.assertEqual(mgr._safe_array_value(np.array(4)), 4) + self.assertEqual(mgr._safe_array_value(np.array([1, 2])).__class__, list) + self.assertIsNone(mgr._safe_array_value(np.array([]))) + + def test_safe_loss_and_prediction_normalization(self): + mgr = LedgeredDataFrameManager(enable_flushing_threads=False, enable_h5_persistence=False) + losses = {"main": np.array([0.1, 0.2]), "aux": np.array([[1.0, 2.0], [3.0, 4.0]])} + out = mgr._safe_loss_dict(losses, idx=1) + self.assertEqual(out["main"], 0.2) + self.assertEqual(out["aux"], [3.0, 4.0]) + + preds_raw = np.array([[[[0.0, 1.0], [2.0, 3.0]]]], dtype=np.float32) + norm = mgr._normalize_preds_raw_uint16(preds_raw) + self.assertEqual(norm.dtype, np.uint16) + self.assertEqual(norm.shape, preds_raw.shape) + self.assertEqual(int(norm.min()), 0) + self.assertEqual(int(norm.max()), 65535) + + passthrough = mgr._normalize_preds_raw_uint16(np.array([1, 2, 3])) + np.testing.assert_array_equal(passthrough, np.array([1, 2, 3])) + + def test_enqueue_batch_buffers_records(self): + mgr = LedgeredDataFrameManager(enable_flushing_threads=False, enable_h5_persistence=False) + + with patch.object(mgr, "flush_async") as flush_async: + mgr.enqueue_batch( + sample_ids=["10", "11"], + preds_raw=np.random.rand(2, 1, 3, 3).astype(np.float32), + preds=np.array([1, 0]), + losses={"loss": np.array([0.4, 0.7])}, + targets=np.array([1, 2]), + step=4, + ) + + self.assertTrue(flush_async.called) + self.assertEqual(len(mgr._buffer), 2) + self.assertIn("sample_id", mgr._buffer["10"]) + self.assertEqual(mgr._buffer["10"][SampleStats.Ex.LAST_SEEN.value], 4) + + def test_enqueue_batch_scalar_target_and_pred_stay_scalar(self): + """A classification-style (B,) target/pred batch must index down to a + true per-sample scalar, not a length-1 array -- a stray trailing axis + upstream previously made this land as [6] instead of 6 (see + expand_dim in src.py's save_signals, since removed).""" + mgr = LedgeredDataFrameManager(enable_flushing_threads=False, enable_h5_persistence=False) + + with patch.object(mgr, "flush_async"): + mgr.enqueue_batch( + sample_ids=["10", "11"], + preds_raw=None, + preds=np.array([1, 0]), + losses=None, + targets=np.array([6, 2]), + step=4, + ) + + target = mgr._buffer["10"][SampleStats.Ex.TARGET.value] + pred = mgr._buffer["10"][SampleStats.Ex.PREDICTION.value] + self.assertEqual(np.asarray(target).ndim, 0) + self.assertEqual(np.asarray(pred).ndim, 0) + self.assertEqual(int(target), 6) + self.assertEqual(int(pred), 1) + + + def test_enqueue_instance_batch_buffers_records(self): + """enqueue_instance_batch should enqueue per-instance records into the SAME + buffer (keyed by (sample_id, annotation_id)) without touching the df.""" + mgr = LedgeredDataFrameManager(enable_flushing_threads=False, enable_h5_persistence=False) + + with patch.object(mgr, "flush_async") as flush_async: + # Instances live at annotation_id >= 1 (instance_id 0 is the sample row). + mgr.enqueue_instance_batch( + sample_ids=["7", "7", "9"], + annotation_ids=[1, 2, 1], + losses={"signal:bbox_loss": np.array([0.1, 0.2, 0.3])}, + targets=[np.array([1, 2, 3, 4]), np.array([5, 6, 7, 8]), np.array([9, 10, 11, 12])], + step=3, + ) + + self.assertTrue(flush_async.called) + # Buffered under composite keys, NOT mutating the dataframe yet. + self.assertEqual(len(mgr._buffer), 3) + self.assertIn(("7", 1), mgr._buffer) + self.assertIn(("7", 2), mgr._buffer) + self.assertIn(("9", 1), mgr._buffer) + rec = mgr._buffer[("7", 2)] + self.assertEqual(rec[SampleStats.Ex.INSTANCE_ID.value], 2) + self.assertEqual(rec["signal:bbox_loss"], 0.2) + self.assertEqual(rec[SampleStats.Ex.LAST_SEEN.value], 3) + self.assertIn(SampleStats.Ex.TARGET.value, rec) + self.assertTrue(mgr._df.empty) # df untouched until flush + + def test_flush_applies_instance_records(self): + """Flushing instance records writes per-(sample_id, annotation_id) values.""" + mgr = LedgeredDataFrameManager(enable_flushing_threads=False, enable_h5_persistence=False) + + # 3-instance sample → 4 rows: instance_id 0 (sample) + 1,2,3 (instances). + target = [np.array([10, 20, 30, 40]), np.array([50, 60, 70, 80]), np.array([90, 100, 110, 120])] + df = pd.DataFrame([{"sample_id": 1, "origin": "train", SampleStats.Ex.TARGET.value: target}]).set_index("sample_id") + mgr.upsert_df(df, origin="train") + + mgr.enqueue_instance_batch( + sample_ids=["1", "1", "1"], + annotation_ids=[1, 2, 3], + losses={"signal:il": np.array([0.5, 0.6, 0.7])}, + step=2, + ) + mgr.flush() + + result = mgr.get_df_view() + self.assertEqual(len(result), 4) # sample row (0) + 3 instance rows + self.assertAlmostEqual(float(result.loc[("1", 1), "signal:il"]), 0.5) + self.assertAlmostEqual(float(result.loc[("1", 2), "signal:il"]), 0.6) + self.assertAlmostEqual(float(result.loc[("1", 3), "signal:il"]), 0.7) + + def test_mixed_sample_and_instance_buffer_flush(self): + """The flush must apply BOTH per-sample (instance_id 0) and per-instance records.""" + mgr = LedgeredDataFrameManager(enable_flushing_threads=False, enable_h5_persistence=False) + + # 2-instance sample → 3 rows: instance_id 0 (sample) + 1,2 (instances). + target = [np.array([1, 2, 3, 4]), np.array([5, 6, 7, 8])] + df = pd.DataFrame([{"sample_id": 1, "origin": "train", SampleStats.Ex.TARGET.value: target}]).set_index("sample_id") + mgr.upsert_df(df, origin="train") + + # Per-sample signal → lands on the canonical sample row (instance_id 0). + mgr.enqueue_batch( + sample_ids=["1"], preds_raw=None, preds=None, + losses={"loss": np.array([0.9])}, step=5, + ) + # Per-instance signal → one value per instance row (instance_id >= 1). + mgr.enqueue_instance_batch( + sample_ids=["1", "1"], annotation_ids=[1, 2], + losses={"signal:il": np.array([0.2, 0.8])}, step=5, + ) + mgr.flush() + + result = mgr.get_df_view() + self.assertEqual(len(result), 3) + # Per-sample value on instance_id 0 only (not broadcast to instance rows). + self.assertAlmostEqual(float(result.loc[("1", 0), "loss"]), 0.9) + self.assertTrue(pd.isna(result.loc[("1", 1), "loss"])) + # Per-instance values on their specific instance rows. + self.assertAlmostEqual(float(result.loc[("1", 1), "signal:il"]), 0.2) + self.assertAlmostEqual(float(result.loc[("1", 2), "signal:il"]), 0.8) + + def test_get_combined_df_surfaces_buffered_instance(self): + """Buffered (unflushed) per-instance values are visible via get_combined_df.""" + mgr = LedgeredDataFrameManager(enable_flushing_threads=False, enable_h5_persistence=False) + + target = [np.array([1, 2, 3, 4]), np.array([5, 6, 7, 8])] + df = pd.DataFrame([{"sample_id": 1, "origin": "train", SampleStats.Ex.TARGET.value: target}]).set_index("sample_id") + mgr.upsert_df(df, origin="train") + + with patch.object(mgr, "flush_async"): + mgr.enqueue_instance_batch( + sample_ids=["1", "1"], annotation_ids=[1, 2], + losses={"signal:il": np.array([0.11, 0.22])}, + ) + # Still buffered (flush_async patched out), but should be merged into the view. + combined = mgr.get_combined_df() + self.assertAlmostEqual(float(combined.loc[("1", 1), "signal:il"]), 0.11) + self.assertAlmostEqual(float(combined.loc[("1", 2), "signal:il"]), 0.22) + + def test_multi_instance_expansion(self): + mgr = LedgeredDataFrameManager(enable_flushing_threads=False, enable_h5_persistence=False) + + # Single sample with 3 instances (detections/annotations) + # Use list of arrays to indicate multiple instances + target = [ + np.array([10, 20, 30, 40]), # instance 0 + np.array([50, 60, 70, 80]), # instance 1 + np.array([90, 100, 110, 120]) # instance 2 + ] + df = pd.DataFrame([{ + "sample_id": 1, + "origin": "train", + SampleStats.Ex.TARGET.value: target, + "metadata": "scene_urban", + "brightness": 0.75 + }]).set_index("sample_id") + + mgr.upsert_df(df, origin="train") + + # 3 instances → 4 rows: instance_id 0 (sample row) + 1,2,3 (the instances). + result_df = mgr.get_df_view() + self.assertEqual(len(result_df), 4) + + # Check multi-index structure + self.assertTrue(isinstance(result_df.index, pd.MultiIndex)) + self.assertEqual(result_df.index.nlevels, 2) + self.assertEqual(result_df.index.names, ['sample_id', 'annotation_id']) + + # Check that all rows have the same sample_id + sample_ids = result_df.index.get_level_values(0) + self.assertTrue((sample_ids == "1").all()) + + # annotation_ids: 0 (sample) then 1, 2, 3 (instances) + annotation_ids = result_df.index.get_level_values(1) + np.testing.assert_array_equal(annotation_ids, [0, 1, 2, 3]) + + # Sample-level metadata lives ONLY on the sample row (instance_id 0); + # instance rows (1..N) carry only their target, everything else empty. + self.assertEqual(result_df.loc[("1", 0), "metadata"], "scene_urban") + self.assertEqual(result_df.loc[("1", 0), "brightness"], 0.75) + for k in (1, 2, 3): + self.assertTrue(pd.isna(result_df.loc[("1", k), "metadata"])) + self.assertTrue(pd.isna(result_df.loc[("1", k), "brightness"])) + # ...but the instance's target IS present. + self.assertIsNotNone(result_df.loc[("1", k), SampleStats.Ex.TARGET.value]) + + def test_multi_instance_different_counts(self): + mgr = LedgeredDataFrameManager(enable_flushing_threads=False, enable_h5_persistence=False) + + # Sample 1: 2 instances (list of arrays) + target1 = [np.array([10, 20, 30, 40]), np.array([50, 60, 70, 80])] + # Sample 2: 3 instances (list of arrays) + target2 = [np.array([1, 2, 3, 4]), np.array([5, 6, 7, 8]), np.array([9, 10, 11, 12])] + # Sample 3: 1 instance (single array) + target3 = np.array([100, 200, 300, 400]) + + df = pd.DataFrame([ + {"sample_id": 1, "origin": "train", SampleStats.Ex.TARGET.value: target1, "metadata": "sample1"}, + {"sample_id": 2, "origin": "train", SampleStats.Ex.TARGET.value: target2, "metadata": "sample2"}, + {"sample_id": 3, "origin": "train", SampleStats.Ex.TARGET.value: target3, "metadata": "sample3"}, + ]).set_index("sample_id") + + mgr.upsert_df(df, origin="train") + + result_df = mgr.get_df_view() + + # Multi-instance samples get a sample row (instance_id 0) + one row per instance; + # a single-array target is the sample's own target (instance_id 0 only). + # Total rows: (1+2) + (1+3) + 1 = 3 + 4 + 1 = 8 + self.assertEqual(len(result_df), 8) + + # Sample 1 (2 instances): instance_id 0 (sample) + 1, 2 + sample1 = result_df.loc["1"] + self.assertEqual(len(sample1), 3) + np.testing.assert_array_equal(sample1.index.tolist(), [0, 1, 2]) + + # Sample 2 (3 instances): instance_id 0 + 1, 2, 3 + sample2 = result_df.loc["2"] + self.assertEqual(len(sample2), 4) + np.testing.assert_array_equal(sample2.index.tolist(), [0, 1, 2, 3]) + + # Sample 3 (single-array target): only the sample row (instance_id 0) + sample3 = result_df.loc["3"] + if isinstance(sample3, pd.Series): + self.assertEqual(sample3.name, 0) + else: + self.assertEqual(len(sample3), 1) + np.testing.assert_array_equal(sample3.index.tolist(), [0]) + + + def test_categorical_memory_optimization(self): + """Test that repetitive string columns are converted to categorical dtype.""" + mgr = LedgeredDataFrameManager(enable_flushing_threads=False, enable_h5_persistence=False) + + # Create data with repetitive 'origin' and 'metadata' columns + # 100 rows but only 3 unique origins and 5 unique metadata values + df = pd.DataFrame([ + { + "sample_id": i, + "origin": ["train", "test", "val"][i % 3], + "metadata": ["urban", "highway", "rural", "city", "suburban"][i % 5], + "loss": float(i) * 0.1, + } + for i in range(100) + ]).set_index("sample_id") + + mgr.upsert_df(df, origin="train") + result_df = mgr.get_df_view() + + # Check that origin column was converted to categorical + # (it's in the categorical_candidates list) + self.assertEqual(result_df["origin"].dtype.name, "category") + self.assertEqual(len(result_df["origin"].cat.categories), 3) + + # Note: metadata is not optimized by default + # (only origin, task_type, and tag columns are optimized automatically) + self.assertEqual(result_df["metadata"].dtype, 'object') + + # Verify data integrity (categorical still works correctly) + self.assertTrue((result_df["origin"] == "train").sum() > 0) + self.assertTrue((result_df["metadata"] == "urban").sum() > 0) + + # Memory usage comparison + # original_bytes = 100 * (len("train") + len("urban")) # Rough estimate + # With categorical: ~100 bytes for codes + ~40 bytes for categories = ~140 bytes + # Real compression achieved by pandas + + + def test_per_sample_buffer_into_multi_index_does_not_corrupt(self): + """Single-level per-sample buffer must not corrupt a multi-index dataframe. + + Regression test for: enqueue_batch produces single-level (sample_id) + records, and _apply_buffer_records used to concat them into a + multi-index dataframe — creating a hybrid index that later crashed + reindex with "cannot reindex on an axis with duplicate labels". + """ + mgr = LedgeredDataFrameManager(enable_flushing_threads=False, enable_h5_persistence=False) + + # Seed multi-instance dataframe (sample 1 has 3 instances, sample 2 has 2) + target1 = [np.array([10, 20, 30, 40]), np.array([50, 60, 70, 80]), np.array([90, 100, 110, 120])] + target2 = [np.array([1, 2, 3, 4]), np.array([5, 6, 7, 8])] + df = pd.DataFrame([ + {"sample_id": 1, "origin": "train", SampleStats.Ex.TARGET.value: target1}, + {"sample_id": 2, "origin": "train", SampleStats.Ex.TARGET.value: target2}, + ]).set_index("sample_id") + mgr.upsert_df(df, origin="train") + self.assertTrue(isinstance(mgr.get_df_view().index, pd.MultiIndex)) + + # Simulate enqueue_batch flushing per-sample signals (single-level keys) + mgr.enqueue_batch( + sample_ids=["1", "2"], + preds_raw=None, + preds=None, + targets=None, + losses={"signals//train/clsf_sample": np.array([0.42, 0.73])}, + step=10, + ) + mgr.flush() + + result = mgr.get_df_view() + + # Index must remain a MultiIndex — no rogue single-level rows + self.assertTrue(isinstance(result.index, pd.MultiIndex)) + self.assertEqual(result.index.nlevels, 2) + self.assertEqual(result.index.names, ["sample_id", "annotation_id"]) + # All index entries must be tuples (no mixed types) + self.assertTrue(all(isinstance(idx, tuple) for idx in result.index)) + + # Per-sample value lands on the canonical sample row (instance_id 0) only. + col = "signals//train/clsf_sample" + self.assertIn(col, result.columns) + self.assertAlmostEqual(result.loc[("1", 0), col], 0.42) + self.assertAlmostEqual(result.loc[("2", 0), col], 0.73) + # Instance rows (>=1) are NOT written by the per-sample path. + self.assertTrue(pd.isna(result.loc[("1", 1), col])) + self.assertTrue(pd.isna(result.loc[("1", 2), col])) + + # Second flush should not crash with "cannot reindex on an axis with duplicate labels" + mgr.enqueue_batch( + sample_ids=["1"], + preds_raw=None, preds=None, targets=None, + losses={"signals//train/clsf_sample": np.array([0.99])}, + step=11, + ) + mgr.flush() # Would raise if bug regressed + result = mgr.get_df_view() + self.assertAlmostEqual(result.loc[("1", 0), col], 0.99) + + def test_normalize_arrays_for_storage_handles_multi_index_row(self): + """_normalize_arrays_for_storage must extract sample_id from MultiIndex tuples. + + Regression test: when the dataframe is multi-indexed, ``row.name`` is a + ``(sample_id, annotation_id)`` tuple. The previous code passed the + tuple directly to ``dataset.get_index_from_sample_id`` which expects a + plain ``sample_id`` — flooding the log with KeyError-string warnings. + """ + mgr = LedgeredDataFrameManager(enable_flushing_threads=False, enable_h5_persistence=False) + + captured = {} + # Fake dataset that records what was passed to it + class _FakeDataset: + def get_index_from_sample_id(self, sid): + captured['sid'] = sid + return 7 + # Stub out the loader lookup so the dataset is reachable + mgr._get_loader_by_origin = lambda origin: type('L', (), {'wrapped_dataset': _FakeDataset()})() + + # Build a row that mimics a multi-index row with an array column + row = pd.Series({ + "origin": "train", + SampleStats.Ex.TARGET.value: np.zeros((30, 30), dtype=np.float32), + }) + row.name = ("12", 0) # MultiIndex-style row.name + + # Should not raise and should pass just the sample_id, not the tuple + mgr._normalize_arrays_for_storage(row) + self.assertEqual(captured.get('sid'), "12") + + def test_enqueue_instance_batch_writes_per_annotation(self): + """enqueue_instance_batch buffers one signal value per (sample_id, annotation_id); + the flush writes them to the correct rows.""" + mgr = LedgeredDataFrameManager(enable_flushing_threads=False, enable_h5_persistence=False) + + # Seed multi-instance dataframe: sample 1 has 3 instances, sample 2 has 2 + target1 = [np.array([10, 20, 30, 40]), np.array([50, 60, 70, 80]), np.array([90, 100, 110, 120])] + target2 = [np.array([1, 2, 3, 4]), np.array([5, 6, 7, 8])] + df = pd.DataFrame([ + {"sample_id": 1, "origin": "train", SampleStats.Ex.TARGET.value: target1}, + {"sample_id": 2, "origin": "train", SampleStats.Ex.TARGET.value: target2}, + ]).set_index("sample_id") + mgr.upsert_df(df, origin="train") + + # Enqueue per-instance IoU signals at instance_id >= 1 (0 = sample row), then flush. + mgr.enqueue_instance_batch( + sample_ids=["1", "1", "1", "2", "2"], + annotation_ids=[1, 2, 3, 1, 2], + losses={"signals//train/iou_instance": np.array([0.9, 0.8, 0.7, 0.5, 0.6])}, + step=5, + ) + mgr.flush() + + result = mgr.get_df_view() + self.assertIn("signals//train/iou_instance", result.columns) + # Each instance has its IoU value at (sample_id, annotation_id >= 1). + self.assertAlmostEqual(result.loc[("1", 1), "signals//train/iou_instance"], 0.9) + self.assertAlmostEqual(result.loc[("1", 2), "signals//train/iou_instance"], 0.8) + self.assertAlmostEqual(result.loc[("1", 3), "signals//train/iou_instance"], 0.7) + self.assertAlmostEqual(result.loc[("2", 1), "signals//train/iou_instance"], 0.5) + self.assertAlmostEqual(result.loc[("2", 2), "signals//train/iou_instance"], 0.6) + + +class TestImportFromStore(unittest.TestCase): + """import_from_store: load another run's persisted per-sample stats (e.g. a + sibling experiment restored from a multi-root viewer) into this ledger.""" + + def setUp(self): + self._tmp = tempfile.TemporaryDirectory(prefix="import_from_store_") + self.tmp_dir = Path(self._tmp.name) + + def tearDown(self): + self._tmp.cleanup() + + def _ledger_with_previous_run(self): + """Three registered samples, still carrying another run's loss and tag.""" + mgr = LedgeredDataFrameManager(enable_flushing_threads=False, enable_h5_persistence=False) + mgr.upsert_df(pd.DataFrame({ + "sample_id": [0, 1, 2], + "annotation_id": [0, 0, 0], + "origin": "train_loader", + SampleStats.Ex.TARGET.value: [7, 2, 1], + SampleStats.Ex.DISCARDED.value: False, + "signals//train-loss-CE": [9.0, 9.0, 9.0], + "signals//margin": [0.5, 0.5, 0.5], + "tag:goldset": [True, True, True], + }).set_index(["sample_id", "annotation_id"]), origin="train_loader") + return mgr + + def _sibling_store(self): + path = self.tmp_dir / "sibling" / "checkpoints" / "data" / "data.h5" + H5DataFrameStore(path).upsert("train_loader", pd.DataFrame({ + "sample_id": [0, 1, 2, 99], # 99: not a sample of this ledger + "annotation_id": [0, 0, 0, 0], + SampleStats.Ex.TARGET.value: [7, 2, 1, 4], + SampleStats.Ex.PREDICTION.value: [7, 3, 1, 4], + SampleStats.Ex.DISCARDED.value: [False, True, False, False], + SampleStats.Ex.NB_SEEN.value: [8, 0, 8, 8], + "signals//train-loss-CE": [0.1, np.nan, 2.5, 0.3], + }).set_index(["sample_id", "annotation_id"])) + return path + + def test_replaces_stats_and_clears_the_previous_runs_columns(self): + mgr = self._ledger_with_previous_run() + path = self._sibling_store() + before = hashlib.md5(path.read_bytes()).hexdigest() + + self.assertEqual(mgr.import_from_store(path), 3, "only samples this ledger knows are imported") + + df = mgr.get_df_view() + self.assertNotIn("99", df.index.get_level_values(0)) + self.assertAlmostEqual(df.loc[("0", 0), "signals//train-loss-CE"], 0.1) + self.assertTrue(pd.isna(df.loc[("1", 0), "signals//train-loss-CE"]), + "a sample the run never scored must not keep the previous run's loss") + self.assertAlmostEqual(df.loc[("2", 0), "signals//train-loss-CE"], 2.5) + self.assertEqual(int(df.loc[("2", 0), SampleStats.Ex.NB_SEEN.value]), 8) + self.assertEqual(str(df.loc[("1", 0), SampleStats.Ex.PREDICTION.value]), "3") + self.assertTrue(bool(df.loc[("1", 0), SampleStats.Ex.DISCARDED.value])) + self.assertFalse(bool(df.loc[("0", 0), SampleStats.Ex.DISCARDED.value])) + + # Columns of another run, absent from the imported one, no row holds anymore. + self.assertNotIn("signals//margin", df.columns) + self.assertNotIn("tag:goldset", df.columns) + + self.assertEqual(hashlib.md5(path.read_bytes()).hexdigest(), before, "the other store is only read") + + def test_another_runs_columns_stay_for_origins_not_imported(self): + mgr = self._ledger_with_previous_run() + mgr.upsert_df(pd.DataFrame({ + "sample_id": [5], "annotation_id": [0], "origin": "test_loader", + "signals//margin": [0.7], "tag:goldset": [True], + }).set_index(["sample_id", "annotation_id"]), origin="test_loader") + + mgr.import_from_store(self._sibling_store(), origins=["train_loader"]) + + df = mgr.get_df_view() + self.assertTrue(df.loc[[("0", 0), ("1", 0), ("2", 0)], "signals//margin"].isna().all()) + self.assertFalse(df.loc[[("0", 0), ("1", 0), ("2", 0)], "tag:goldset"].astype(bool).any()) + self.assertAlmostEqual(df.loc[("5", 0), "signals//margin"], 0.7) + self.assertTrue(bool(df.loc[("5", 0), "tag:goldset"])) + + def test_missing_store_or_unregistered_origin_imports_nothing(self): + mgr = self._ledger_with_previous_run() + self.assertEqual(mgr.import_from_store(self.tmp_dir / "nowhere" / "data.h5"), 0) + self.assertEqual(mgr.import_from_store(self._sibling_store(), origins=["test_loader"]), 0) + self.assertAlmostEqual(mgr.get_df_view().loc[("0", 0), "signals//train-loss-CE"], 9.0) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/data/test_dataframe_rewind.py b/tests/data/test_dataframe_rewind.py new file mode 100644 index 00000000..6f393b2b --- /dev/null +++ b/tests/data/test_dataframe_rewind.py @@ -0,0 +1,173 @@ +"""Tests for ``LedgeredDataFrameManager.rewind_to_step``. + +Restoring a checkpoint moves the model's age backwards; these cover the ledger +catching up with it — signal values rolled back to the step's history, +``last_seen``/``nb_seen`` recomputed, and predictions from the discarded model +state cleared. +""" + +import unittest + +import numpy as np +import pandas as pd + +from weightslab.data.dataframe_manager import LedgeredDataFrameManager +from weightslab.data.sample_stats import SampleStats + + +LAST_SEEN = SampleStats.Ex.LAST_SEEN.value +NB_SEEN = SampleStats.Ex.NB_SEEN.value +PREDICTION = SampleStats.Ex.PREDICTION.value +PREDICTION_RAW = SampleStats.Ex.PREDICTION_RAW.value +TARGET = SampleStats.Ex.TARGET.value + + +def _manager(): + return LedgeredDataFrameManager(enable_flushing_threads=False, enable_h5_persistence=False) + + +def _row(sample_id, loss, last_seen, nb_seen, **extra): + row = { + "sample_id": sample_id, + "origin": "train", + "signals//loss": loss, + LAST_SEEN: last_seen, + NB_SEEN: nb_seen, + PREDICTION: np.array([1, 2]), + PREDICTION_RAW: np.array([0.3, 0.7]), + TARGET: 7, + } + row.update(extra) + return row + + +def _seed(mgr, rows): + mgr.upsert_df(pd.DataFrame(rows).set_index("sample_id"), origin="train") + return mgr + + +def _cell(mgr, sample_id, column): + return mgr.get_df_view().loc[(sample_id, 0), column] + + +class TestRewindToStep(unittest.TestCase): + def setUp(self): + self.mgr = _manager() + _seed(self.mgr, [ + # Ahead of the restore point: everything about it is stale. + _row("1", loss=0.1, last_seen=9, nb_seen=4), + # Exactly at the restore point: still valid, must not be touched. + _row("2", loss=0.3, last_seen=5, nb_seen=2), + # Never seen. + _row("3", loss=np.nan, last_seen=-1, nb_seen=0), + ]) + self.state = { + "1": {"signals": {"loss": 0.4}, "last_seen": 5, "nb_seen": 3}, + "2": {"signals": {"loss": 0.3}, "last_seen": 5, "nb_seen": 2}, + } + + def test_rolls_signals_back_to_their_value_at_the_step(self): + self.assertEqual(self.mgr.rewind_to_step(5, self.state), 1) + self.assertAlmostEqual(_cell(self.mgr, "1", "signals//loss"), 0.4, places=6) + + def test_recomputes_last_seen_and_nb_seen(self): + self.mgr.rewind_to_step(5, self.state) + self.assertEqual(int(_cell(self.mgr, "1", LAST_SEEN)), 5) + self.assertEqual(int(_cell(self.mgr, "1", NB_SEEN)), 3) + + def test_clears_predictions_from_the_discarded_model_state(self): + self.mgr.rewind_to_step(5, self.state) + self.assertIsNone(_cell(self.mgr, "1", PREDICTION)) + self.assertIsNone(_cell(self.mgr, "1", PREDICTION_RAW)) + + def test_keeps_targets(self): + # Ground truth does not come from the model, so a rewind must not lose it. + self.mgr.rewind_to_step(5, self.state) + self.assertEqual(_cell(self.mgr, "1", TARGET), 7) + + def test_samples_not_ahead_of_the_step_are_untouched(self): + self.mgr.rewind_to_step(5, self.state) + + for sample_id, loss, last_seen, nb_seen in (("2", 0.3, 5, 2), ("3", np.nan, -1, 0)): + value = _cell(self.mgr, sample_id, "signals//loss") + if np.isnan(loss): + self.assertTrue(np.isnan(value)) + else: + self.assertAlmostEqual(value, loss, places=6) + self.assertEqual(int(_cell(self.mgr, sample_id, LAST_SEEN)), last_seen) + self.assertEqual(int(_cell(self.mgr, sample_id, NB_SEEN)), nb_seen) + # Sample 2's prediction is from a step the restored model still owns. + self.assertIsNotNone(_cell(self.mgr, "2", PREDICTION)) + + def test_signal_with_no_history_that_old_becomes_nan(self): + mgr = _seed(_manager(), [ + _row("1", loss=0.1, last_seen=9, nb_seen=4, **{"signals//late": 42.0}), + ]) + # "late" was first recorded after the restore point, so it has no value. + mgr.rewind_to_step(5, {"1": {"signals": {"loss": 0.4}, "last_seen": 5, "nb_seen": 3}}) + + self.assertAlmostEqual(_cell(mgr, "1", "signals//loss"), 0.4, places=6) + self.assertTrue(np.isnan(_cell(mgr, "1", "signals//late"))) + + def test_sample_absent_from_the_history_is_reset_to_never_seen(self): + rewound = self.mgr.rewind_to_step(5, {}) + + self.assertEqual(rewound, 1) # only sample 1 was ahead of the step + self.assertTrue(np.isnan(_cell(self.mgr, "1", "signals//loss"))) + self.assertEqual(int(_cell(self.mgr, "1", LAST_SEEN)), + SampleStats.DEFAULTS[LAST_SEEN]) + self.assertEqual(int(_cell(self.mgr, "1", NB_SEEN)), SampleStats.DEFAULTS[NB_SEEN]) + + def test_reset_predictions_false_keeps_them(self): + self.mgr.rewind_to_step(5, self.state, reset_predictions=False) + + self.assertIsNotNone(_cell(self.mgr, "1", PREDICTION)) + # The rest of the rewind still applies. + self.assertAlmostEqual(_cell(self.mgr, "1", "signals//loss"), 0.4, places=6) + self.assertEqual(int(_cell(self.mgr, "1", LAST_SEEN)), 5) + + def test_nothing_ahead_of_the_step_is_a_noop(self): + self.assertEqual(self.mgr.rewind_to_step(9, self.state), 0) + self.assertAlmostEqual(_cell(self.mgr, "1", "signals//loss"), 0.1, places=6) + self.assertIsNotNone(_cell(self.mgr, "1", PREDICTION)) + + def test_empty_ledger_is_a_noop(self): + self.assertEqual(_manager().rewind_to_step(5, self.state), 0) + + def test_ledger_without_last_seen_is_a_noop(self): + mgr = _manager() + mgr.upsert_df( + pd.DataFrame([{"sample_id": "1", "origin": "train", "signals//loss": 0.1}]) + .set_index("sample_id"), origin="train") + self.assertEqual(mgr.rewind_to_step(5, self.state), 0) + + def test_instance_rows_are_left_alone(self): + # Sample-level values live on annotation 0; instance rows carry their own + # per-instance state, which this rewind does not read. + mgr = _manager() + frame = pd.DataFrame([ + {"sample_id": "1", "annotation_id": 0, "origin": "train", + "signals//loss": 0.1, LAST_SEEN: 9, NB_SEEN: 4, PREDICTION: np.array([1, 2])}, + {"sample_id": "1", "annotation_id": 1, "origin": "train", + "signals//iou": 0.9, LAST_SEEN: 9, NB_SEEN: 4, PREDICTION: np.array([3, 4])}, + ]).set_index(["sample_id", "annotation_id"]) + mgr.upsert_df(frame, origin="train") + + self.assertEqual( + mgr.rewind_to_step(5, {"1": {"signals": {"loss": 0.4}, "last_seen": 5, "nb_seen": 3}}), + 1) + + view = mgr.get_df_view() + self.assertAlmostEqual(view.loc[("1", 0), "signals//loss"], 0.4, places=6) + self.assertIsNone(view.loc[("1", 0), PREDICTION]) + self.assertAlmostEqual(view.loc[("1", 1), "signals//iou"], 0.9, places=6) + self.assertIsNotNone(view.loc[("1", 1), PREDICTION]) + + def test_marks_rewound_samples_dirty_for_persistence(self): + self.mgr.clear_view_dirty() + self.mgr.rewind_to_step(5, self.state) + self.assertIn("1", {str(s) for s in self.mgr.take_view_dirty()}) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/data/test_flush_pipeline.py b/tests/data/test_flush_pipeline.py new file mode 100644 index 00000000..70f0b72a --- /dev/null +++ b/tests/data/test_flush_pipeline.py @@ -0,0 +1,276 @@ +""" +Unit tests for the buffer → DataFrame → H5 flush pipeline in LedgeredDataFrameManager. + +Verified behaviors: + 1. flush() releases _buffer_lock before DF/H5 work — training can enqueue during a flush. + 2. flush_async() returns after buffer drain, not after H5 write completes. + 3. If buffer refills while H5 write is ongoing, training waits only until the flush + thread drains the buffer again (not until H5 finishes). + 4. In-memory buffer is bounded to ≤ flush_max_rows records at any point. +""" + +import time +import threading +import unittest +import numpy as np + +from unittest.mock import patch + +from weightslab.data.dataframe_manager import LedgeredDataFrameManager + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + +def _make_mgr(flush_max_rows=4, enable_flushing_threads=False) -> LedgeredDataFrameManager: + mgr = LedgeredDataFrameManager( + flush_interval=60.0, # disable periodic timer during tests + flush_max_rows=flush_max_rows, + enable_flushing_threads=enable_flushing_threads, + enable_h5_persistence=False, + ) + return mgr + + +def _enqueue(mgr, sample_ids, step=0): + n = len(sample_ids) + mgr.enqueue_batch( + sample_ids=sample_ids, + preds_raw=np.random.rand(n, 1, 4, 4).astype(np.float32), + preds=np.zeros(n, dtype=np.int64), + losses={"loss": np.ones(n, dtype=np.float32) * 0.1}, + targets=np.zeros(n, dtype=np.int64), + step=step, + ) + + +# --------------------------------------------------------------------------- +# Test 1: flush() releases _buffer_lock before DF/H5 work +# --------------------------------------------------------------------------- + +class TestFlushReleasesBufferLockEarly(unittest.TestCase): + """flush() must release _buffer_lock right after draining so that + enqueue_batch() is not blocked during the (slow) DF/H5 phase.""" + + def test_enqueue_succeeds_while_flush_doing_df_write(self): + mgr = _make_mgr(enable_flushing_threads=False) + + # Seed the DataFrame so _apply_buffer_records has rows to update. + for i in range(4): + mgr._buffer[str(i)] = {"sample_id": str(i), "origin": "train"} + mgr._drain_buffer() + + enqueue_started = threading.Event() + enqueue_finished = threading.Event() + df_write_started = threading.Event() + df_write_may_proceed = threading.Event() + + original_apply = mgr._apply_buffer_records + + def slow_apply(records): + df_write_started.set() + df_write_may_proceed.wait(timeout=5) + original_apply(records) + + def do_flush(): + with patch.object(mgr, "_apply_buffer_records", side_effect=slow_apply): + # Pre-fill buffer so flush has something to drain. + with mgr._buffer_lock: + for i in range(4): + mgr._buffer[str(i)] = {"sample_id": str(i), "origin": "train"} + mgr.flush() + + def do_enqueue(): + enqueue_started.set() + _enqueue(mgr, ["99", "100"]) + enqueue_finished.set() + + flush_thread = threading.Thread(target=do_flush) + flush_thread.start() + + # Wait until flush is inside the slow DF write (buffer already drained). + df_write_started.wait(timeout=5) + + # Now enqueue from a second thread — must NOT be blocked on _buffer_lock. + enqueue_thread = threading.Thread(target=do_enqueue) + enqueue_thread.start() + + # Give enqueue thread 1 second to complete; if it can't, flush is holding + # _buffer_lock too long. + enqueue_finished.wait(timeout=1.0) + self.assertTrue( + enqueue_finished.is_set(), + "enqueue_batch() was blocked while flush() was doing DF work — " + "_buffer_lock was held too long.", + ) + + df_write_may_proceed.set() + flush_thread.join(timeout=5) + enqueue_thread.join(timeout=5) + + +# --------------------------------------------------------------------------- +# Test 2: flush_async() returns after buffer drain, not after H5 write +# --------------------------------------------------------------------------- + +class TestFlushAsyncReturnsAfterBufferDrain(unittest.TestCase): + """flush_async() must return as soon as the buffer has been drained — + not after the (potentially slow) H5 write. + + Scenario: fill buffer once. flush_async() (called inside enqueue_batch) should + return after the flush thread drains the buffer (~thread wake-up time), well + before the slow H5 write completes. + """ + + def test_flush_async_does_not_wait_for_h5(self): + H5_WRITE_DELAY = 1.0 # seconds — intentionally slow + + mgr = _make_mgr(flush_max_rows=4, enable_flushing_threads=True) + + original_flush_to_h5 = mgr._flush_to_h5_if_needed + + def slow_h5(*args, **kwargs): + time.sleep(H5_WRITE_DELAY) + original_flush_to_h5(*args, **kwargs) + + with patch.object(mgr, "_flush_to_h5_if_needed", side_effect=slow_h5): + # Fill buffer to capacity — this triggers flush_async() inside enqueue_batch. + # flush_async() should return after buffer drain (~thread wake-up, <<1s), + # NOT after the H5 write (1s). + t0 = time.time() + _enqueue(mgr, ["0", "1", "2", "3"]) + elapsed = time.time() - t0 + + mgr.stop() + + # The flush thread needs a few hundred ms to wake up and drain. + # Anything under 80% of H5_WRITE_DELAY proves we did NOT wait for H5. + self.assertLess( + elapsed, + H5_WRITE_DELAY * 0.8, + f"flush_async() waited {elapsed:.2f}s — it should return after buffer " + f"drain (~ms), not after the {H5_WRITE_DELAY}s H5 write.", + ) + + +# --------------------------------------------------------------------------- +# Test 3: buffer refills while H5 writing — training waits, then resumes +# --------------------------------------------------------------------------- + +class TestBufferRefillDuringH5Write(unittest.TestCase): + """If the buffer fills while an H5 write is in progress, the training + thread must wait (bounded wait) and resume once the flush thread drains + the buffer again in the next cycle.""" + + def test_training_resumes_after_second_drain(self): + FLUSH_MAX = 4 + H5_WRITE_DELAY = 0.3 # seconds + + mgr = _make_mgr(flush_max_rows=FLUSH_MAX, enable_flushing_threads=True) + + flush_cycle_count = {"n": 0} + original_flush_to_h5 = mgr._flush_to_h5_if_needed + + def counting_slow_h5(*args, **kwargs): + flush_cycle_count["n"] += 1 + time.sleep(H5_WRITE_DELAY) + original_flush_to_h5(*args, **kwargs) + + second_enqueue_returned = threading.Event() + + def training_sim(): + _enqueue(mgr, [str(i) for i in range(FLUSH_MAX)]) # fills buffer, triggers flush + time.sleep(0.05) # let flush thread start H5 write + _enqueue(mgr, [str(i) for i in range(FLUSH_MAX, FLUSH_MAX * 2)]) # refill + second_enqueue_returned.set() + + with patch.object(mgr, "_flush_to_h5_if_needed", side_effect=counting_slow_h5): + t = threading.Thread(target=training_sim) + t.start() + completed = second_enqueue_returned.wait(timeout=H5_WRITE_DELAY * 6) + + mgr.stop() + t.join(timeout=5) + + self.assertTrue( + completed, + "Training thread never resumed after the second buffer fill.", + ) + self.assertGreaterEqual( + flush_cycle_count["n"], 1, + "Flush thread never ran an H5 write cycle.", + ) + + +# --------------------------------------------------------------------------- +# Test 4: memory stays bounded at <= flush_max_rows during concurrent load +# --------------------------------------------------------------------------- + +class TestBufferMemoryBound(unittest.TestCase): + """The buffer must never hold more than flush_max_rows records because + flush_async() blocks the training thread when the buffer is at capacity.""" + + def test_buffer_never_exceeds_max_rows(self): + FLUSH_MAX = 8 + TOTAL_SAMPLES = 200 + + mgr = _make_mgr(flush_max_rows=FLUSH_MAX, enable_flushing_threads=True) + + max_buffer_seen = {"n": 0} + + # Track high-water mark inside the buffer lock. + observation_lock = threading.Lock() + + original_apply = mgr._apply_buffer_records + + def tracking_apply(records): + # Snapshot buffer size right before drain completes. + with mgr._buffer_lock: + with observation_lock: + max_buffer_seen["n"] = max(max_buffer_seen["n"], len(mgr._buffer)) + original_apply(records) + + with patch.object(mgr, "_apply_buffer_records", side_effect=tracking_apply): + for batch_start in range(0, TOTAL_SAMPLES, FLUSH_MAX // 2): + ids = [str(i) for i in range(batch_start, batch_start + FLUSH_MAX // 2)] + _enqueue(mgr, ids, step=batch_start) + with mgr._buffer_lock: + current = len(mgr._buffer) + with observation_lock: + max_buffer_seen["n"] = max(max_buffer_seen["n"], current) + + mgr.stop() + + self.assertLessEqual( + max_buffer_seen["n"], + FLUSH_MAX, + f"Buffer reached {max_buffer_seen['n']} records — exceeded flush_max_rows={FLUSH_MAX}.", + ) + + +# --------------------------------------------------------------------------- +# Test 5: flush() uses blocking _apply_buffer_records (not nonblocking) +# --------------------------------------------------------------------------- + +class TestFlushUsesBlockingApply(unittest.TestCase): + """flush() must call _apply_buffer_records (blocking) so that records are + guaranteed to land in the DataFrame even under lock contention.""" + + def test_flush_calls_blocking_apply(self): + mgr = _make_mgr(enable_flushing_threads=False) + + with mgr._buffer_lock: + for i in range(3): + mgr._buffer[str(i)] = {"sample_id": str(i), "origin": "train"} + + with patch.object(mgr, "_apply_buffer_records", wraps=mgr._apply_buffer_records) as blocking_mock, \ + patch.object(mgr, "_apply_buffer_records_nonblocking", wraps=mgr._apply_buffer_records_nonblocking) as nonblocking_mock: + mgr.flush() + + blocking_mock.assert_called_once() + nonblocking_mock.assert_not_called() + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/data/test_h5_array_store.py b/tests/data/test_h5_array_store.py new file mode 100644 index 00000000..bf1add0b --- /dev/null +++ b/tests/data/test_h5_array_store.py @@ -0,0 +1,216 @@ +import shutil +import tempfile +import unittest +from pathlib import Path + +import numpy as np +import pandas as pd +import h5py + +from weightslab.data.h5_array_store import H5ArrayStore +from weightslab.data.array_proxy import ArrayH5Proxy, convert_dataframe_to_proxies + + +# --------------------------------------------------------------------------- +# Existing functional tests +# --------------------------------------------------------------------------- + +class TestH5ArrayStore(unittest.TestCase): + def setUp(self): + self.tmpdir = tempfile.mkdtemp() + self.array_path = Path(self.tmpdir) / "arrays.h5" + self.store = H5ArrayStore(self.array_path, auto_normalize=True) + + def tearDown(self): + shutil.rmtree(self.tmpdir, ignore_errors=True) + + def test_save_and_load_preserve_original(self): + arr = np.arange(12, dtype=np.float32).reshape(3, 4) + + path_ref = self.store.save_array(1, "prediction", arr, preserve_original=True) + + self.assertIsNotNone(path_ref) + self.assertTrue(str(path_ref).endswith(":/1/prediction")) + + loaded = self.store.load_array(path_ref) + self.assertIsNotNone(loaded) + np.testing.assert_array_equal(loaded, arr) + self.assertEqual(loaded.dtype, arr.dtype) + + def test_batch_save_and_load(self): + arrays_dict = { + '10': { + "prediction": np.ones((2, 2), dtype=np.float32), + "target": np.zeros((2, 2), dtype=np.int32), + }, + '11': { + "prediction": np.full((3,), 7, dtype=np.int16), + }, + } + + refs = self.store.save_arrays_batch(arrays_dict, preserve_original=True) + self.assertEqual(set(refs.keys()), {'10', '11'}) + self.assertIn("prediction", refs['10']) + + loaded = self.store.load_arrays_batch(refs) + self.assertEqual(set(loaded.keys()), {'10', '11'}) + np.testing.assert_array_equal(loaded['10']["prediction"], arrays_dict['10']["prediction"]) + np.testing.assert_array_equal(loaded['10']["target"], arrays_dict['10']["target"]) + np.testing.assert_array_equal(loaded['11']["prediction"], arrays_dict['11']["prediction"]) + + def test_delete_sample(self): + arr = np.ones((3, 3), dtype=np.float32) + path_ref = self.store.save_array(99, "prediction", arr, preserve_original=True) + self.assertTrue(self.store.delete_sample(99)) + + # After deletion, load should return None + self.assertIsNone(self.store.load_array(path_ref)) + # File should still exist + self.assertTrue(self.store.get_path().exists()) + + def test_convert_dataframe_autoload_partial(self): + # Save two arrays and build a dataframe with path references + pred = np.random.rand(2, 2).astype(np.float32) + tgt = np.random.randint(0, 3, size=(2, 2)).astype(np.int32) + + pred_ref = self.store.save_array(5, "prediction", pred, preserve_original=True) + tgt_ref = self.store.save_array(5, "target", tgt, preserve_original=True) + + df = pd.DataFrame( + {"prediction": [pred_ref], "target": [tgt_ref]}, + index=pd.MultiIndex.from_arrays([[5], [0]], names=["sample_id", "annotation_id"]), + ) + + # Autoload only prediction; target stays proxy + df_out = convert_dataframe_to_proxies( + df, + array_columns=["prediction", "target"], + array_store=self.store, + autoload=["prediction"], + return_proxies=True, + ) + + self.assertIsInstance(df_out.loc[(5, 0), "prediction"], np.ndarray) + self.assertIsInstance(df_out.loc[(5, 0), "target"], ArrayH5Proxy) + # Accessing the proxy should load the array transparently + np.testing.assert_array_equal(df_out.loc[(5, 0), "target"], tgt) + + +# --------------------------------------------------------------------------- +# Crash-safety tests +# --------------------------------------------------------------------------- + +class TestH5ArrayStoreCrashSafety(unittest.TestCase): + """ + Verifies that arrays.h5 remains readable after a crash during write. + + Strategy: simulate crash scenarios by directly manipulating files + (creating temp files, backups) and verifying that recover() handles them. + """ + + def setUp(self): + self.tmpdir = tempfile.mkdtemp() + self.array_path = Path(self.tmpdir) / "arrays.h5" + + def tearDown(self): + shutil.rmtree(self.tmpdir, ignore_errors=True) + + def _make_store(self) -> H5ArrayStore: + return H5ArrayStore(self.array_path, auto_normalize=False) + + def _populate(self) -> dict: + """Write sample_id=1 into arrays.h5 and return path refs.""" + store = self._make_store() + initial = {'1': {"prediction": np.ones((8, 8), dtype=np.uint8)}} + refs = store.save_arrays_batch(initial, preserve_original=True) + self.assertTrue(self.array_path.exists()) + return refs + + def test_kill_phase1_main_file_untouched(self): + """ + Leftover temp file from phase 1 (before backup created) must leave + arrays.h5 untouched. recover() cleans up the dangling temp file. + """ + refs = self._populate() + + # Simulate a crash during phase 1 by creating a leftover temp file + temp_file = self.array_path.with_suffix(".h5.writing_abc12345") + with h5py.File(str(temp_file), 'w') as f: + f.create_group('2') + + # No backup should exist (phase 2 never started) + self.assertFalse(self.array_path.with_suffix(".h5.backup").exists()) + + # recover() removes the temp file and leaves arrays.h5 intact + fresh = self._make_store() + fresh.recover() + self.assertEqual( + list(self.array_path.parent.glob("arrays.h5.writing_*")), + [], + "recover() should have deleted the temp file", + ) + + # Original data is fully readable + loaded = fresh.load_arrays_batch(refs) + self.assertIn('1', loaded) + np.testing.assert_array_equal( + loaded['1']["prediction"], + np.ones((8, 8), dtype=np.uint8), + ) + + def test_kill_phase2_recover_restores_backup(self): + """ + Leftover backup file from phase 2 (after backup created but before merge + completed) must be restored by recover(). + """ + refs = self._populate() + + # Simulate a crash during phase 2 by creating a backup + backup = self.array_path.with_suffix(".h5.backup") + shutil.copy2(self.array_path, backup) + + # Corrupt the main file to simulate incomplete merge + with h5py.File(str(self.array_path), 'a') as f: + if '2' not in f: + f.create_group('2') + + # Verify backup exists + self.assertTrue(backup.exists()) + + # recover() restores the backup and removes it + fresh = self._make_store() + fresh.recover() + self.assertFalse(backup.exists(), "recover() must remove the backup after restoring") + + # Original data is readable after restore + loaded = fresh.load_arrays_batch(refs) + self.assertIn('1', loaded) + np.testing.assert_array_equal( + loaded['1']["prediction"], + np.ones((8, 8), dtype=np.uint8), + ) + + # The batch that was being written when crashed must not appear + self.assertIsNone( + fresh.load_array("arrays.h5:/2/target"), + "Incomplete batch data must not be present after recover()", + ) + + def test_clean_write_leaves_no_temp_or_backup(self): + """After a normal successful write, no temp or backup files are left.""" + store = self._make_store() + store.save_arrays_batch( + {1: {"prediction": np.zeros((4, 4), dtype=np.uint8)}}, + preserve_original=True, + ) + self.assertFalse(self.array_path.with_suffix(".h5.backup").exists()) + self.assertEqual(list(self.array_path.parent.glob("arrays.h5.writing_*")), []) + + def test_recover_safe_on_empty_directory(self): + """recover() must not raise when arrays.h5 does not exist yet.""" + store = self._make_store() + store.recover() # Should complete without error + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/data/test_h5_dataframe_store.py b/tests/data/test_h5_dataframe_store.py new file mode 100644 index 00000000..9282d442 --- /dev/null +++ b/tests/data/test_h5_dataframe_store.py @@ -0,0 +1,243 @@ +import shutil +import tempfile +import unittest +from pathlib import Path + +import pandas as pd +import numpy as np + +from weightslab.data.h5_dataframe_store import H5DataFrameStore +from weightslab.data.sample_stats import SampleStatsEx + + +class TestH5DataFrameStore(unittest.TestCase): + def setUp(self): + self.tmpdir = tempfile.mkdtemp() + self.h5_path = Path(self.tmpdir) / "data_with_ops.h5" + self.store = H5DataFrameStore(self.h5_path) + + def tearDown(self): + shutil.rmtree(self.tmpdir, ignore_errors=True) + + def test_destructive_upsert_takes_a_backup(self): + """The rewrite path must actually back the file up first. + + The copy used to run while the store was still open; on Windows HDF5 + locks an open file, so the copy failed with "Permission denied" and + every rewrite ran without a backup. + """ + self.store.upsert("train", pd.DataFrame({"sample_id": [0, 1], "loss": [0.5, 0.7]}).set_index("sample_id")) + + results = [] + original = self.store._create_backup + + def spy(): + results.append(original()) + return results[-1] + + self.store._create_backup = spy + # New rows -> not an in-place update -> read-merge-rewrite with a backup. + self.store.upsert("train", pd.DataFrame({"sample_id": [2, 3], "loss": [0.1, 0.2]}).set_index("sample_id")) + + self.assertEqual(len(results), 1) + self.assertIsNotNone(results[0], "backup copy failed (store still open while copying?)") + self.assertEqual(len(self.store.load("train")), 4) + + def test_upsert_and_load_all(self): + """Original test: single-level index backward compatibility.""" + train_df = pd.DataFrame( + { + "sample_id": [1, 2], + f"{SampleStatsEx.TAG.value}:a": 1, + SampleStatsEx.DISCARDED.value: [False, True], + } + ).set_index("sample_id") + eval_df = pd.DataFrame( + { + "sample_id": [3], + + f"{SampleStatsEx.TAG.value}:b": 1, + SampleStatsEx.DISCARDED.value: [False], + } + ).set_index("sample_id") + + self.store.upsert("train", train_df) + self.store.upsert("eval", eval_df) + + loaded = self.store.load_all(["train", "eval"]) + self.assertEqual(set(loaded["origin"].unique()), {"train", "eval"}) + self.assertIn("sample_id", loaded.columns) + self.assertEqual(len(loaded), 3) + # Ensure values are preserved (re-index on both levels — every row is multi-indexed). + train_rows = loaded[loaded["origin"] == "train"].set_index(["sample_id", "annotation_id"]) + self.assertTrue(train_rows.loc[(2, 0), SampleStatsEx.DISCARDED.value]) + + def test_single_level_input_is_promoted_to_multi_index(self): + """A bare single-level (sample_id) frame must be PROMOTED to the + (sample_id, annotation_id=0) multi-index on write — no dataframe is ever + persisted/loaded with only sample_id as its index.""" + # Create single-level indexed dataframe (legacy / convenience input) + df = pd.DataFrame({ + 'sample_id': [10, 11, 12], + 'brightness': [0.75, 0.82, 0.65], + 'discarded': [False, False, True] + }).set_index('sample_id') + + # Write + self.store.upsert('train', df) + + # Read + loaded = self.store.load('train') + + # The store restores both index levels as columns; annotation_id must exist + # and be all-zero (the canonical sample rows). + self.assertIn('sample_id', loaded.columns) + self.assertIn('annotation_id', loaded.columns) + self.assertEqual(list(loaded['annotation_id']), [0, 0, 0]) + self.assertEqual(len(loaded), 3) + self.assertEqual(sorted(loaded['sample_id'].astype(int)), [10, 11, 12]) + # Re-indexing on both levels yields a proper 2-level MultiIndex. + mi = loaded.set_index(['sample_id', 'annotation_id']) + self.assertIsInstance(mi.index, pd.MultiIndex) + self.assertEqual(mi.index.nlevels, 2) + + def test_multi_index_write_read_round_trip(self): + """Verify multi-index (sample_id, annotation_id) is preserved through write/read.""" + # Create multi-index dataframe (expanded format) + df = pd.DataFrame({ + 'brightness': [0.75, 0.78, 0.82], + 'iou': [0.72, 0.58, 0.89], + }) + df.index = pd.MultiIndex.from_arrays( + [[100, 100, 101], [0, 1, 0]], + names=['sample_id', 'annotation_id'] + ) + + # Write + self.store.upsert('train', df) + + # Read + loaded = self.store.load('train') + + # Verify multi-index is restored + self.assertIn('sample_id', loaded.columns) + self.assertIn('annotation_id', loaded.columns) + self.assertEqual(list(loaded['sample_id']), [100, 100, 101]) + self.assertEqual(list(loaded['annotation_id']), [0, 1, 0]) + # Note: index may or may not be MultiIndex after read, but columns are restored + self.assertEqual(len(loaded), 3) + + def test_categorical_tags_preservation(self): + """Verify categorical tags are preserved through write/read.""" + df = pd.DataFrame({ + 'sample_id': [1, 2, 3], + 'brightness': [0.75, 0.82, 0.65], + 'tag:quality': ['high', 'low', 'high'], # String tag + 'tag:outdoor': [True, False, True], # Boolean tag + }).set_index('sample_id') + + # Write (should optimize to categorical) + self.store.upsert('train', df) + + # Read + loaded = self.store.load('train') + + # Verify categorical dtypes are preserved + # Note: HDF5 with format="table" preserves categorical dtype + self.assertIn('tag:quality', loaded.columns) + self.assertIn('tag:outdoor', loaded.columns) + + # Check if categorical (may be categorical or object depending on HDF5 behavior) + # The important thing is that the values are correct + self.assertEqual(list(loaded['tag:quality']), ['high', 'low', 'high']) + # Boolean tags are preserved (either as bool or converted to string, both acceptable) + outdoor_values = list(loaded['tag:outdoor']) + # Check they are either boolean or string representation + self.assertTrue( + outdoor_values == [True, False, True] or + outdoor_values == ['True', 'False', 'True'] + ) + + def test_categorical_tags_memory_optimization(self): + """Verify categorical optimization reduces memory usage.""" + # Create dataframe with repetitive tag values (many samples) + n_samples = 1000 + df = pd.DataFrame({ + 'sample_id': range(n_samples), + 'brightness': np.random.rand(n_samples), + 'tag:quality': ['high' if i % 2 == 0 else 'low' for i in range(n_samples)], + }).set_index('sample_id') + + # Get memory before categorical optimization + normalized_df = self.store._normalize_for_write(df) + + self.assertIsNotNone(normalized_df) + + def test_multi_index_with_tags(self): + """Verify multi-index and categorical tags work together.""" + df = pd.DataFrame({ + 'brightness': [0.75, 0.78, 0.82], + 'iou': [0.72, 0.58, 0.89], + 'tag:quality': ['high', 'low', 'high'], + 'tag:object': ['person', 'car', 'person'], + }) + df.index = pd.MultiIndex.from_arrays( + [[100, 100, 101], [0, 1, 0]], + names=['sample_id', 'annotation_id'] + ) + + # Write (should preserve multi-index and optimize tags) + self.store.upsert('train', df) + + # Read + loaded = self.store.load('train') + + # Verify both features work together + self.assertIn('sample_id', loaded.columns) + self.assertIn('annotation_id', loaded.columns) + self.assertIn('tag:quality', loaded.columns) + self.assertIn('tag:object', loaded.columns) + + self.assertEqual(list(loaded['sample_id']), [100, 100, 101]) + self.assertEqual(list(loaded['annotation_id']), [0, 1, 0]) + self.assertEqual(list(loaded['tag:quality']), ['high', 'low', 'high']) + + def test_upsert_merge_multi_index(self): + """Verify upsert merge works correctly with multi-index.""" + # Initial data + df1 = pd.DataFrame({ + 'brightness': [0.75, 0.78], + 'iou': [0.72, 0.58], + }) + df1.index = pd.MultiIndex.from_arrays( + [[100, 100], [0, 1]], + names=['sample_id', 'annotation_id'] + ) + + self.store.upsert('train', df1) + + # Update with new data for same sample but different annotation + df2 = pd.DataFrame({ + 'brightness': [0.80], # Update brightness for annotation 1 + 'iou': [0.60], + }) + df2.index = pd.MultiIndex.from_arrays( + [[100], [1]], + names=['sample_id', 'annotation_id'] + ) + + self.store.upsert('train', df2) + + # Read and verify merge worked + loaded = self.store.load('train') + self.assertEqual(len(loaded), 2) + + # Check that annotation 1 was updated + anno1_rows = loaded[loaded['annotation_id'] == 1] + self.assertEqual(len(anno1_rows), 1) + # Value should be from df2 (updated) + self.assertAlmostEqual(anno1_rows['brightness'].iloc[0], 0.80) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/data/test_point_cloud_utils.py b/tests/data/test_point_cloud_utils.py new file mode 100644 index 00000000..62243a9d --- /dev/null +++ b/tests/data/test_point_cloud_utils.py @@ -0,0 +1,330 @@ +"""Tests for point-cloud preview utilities (BEV rendering / box projection / +binary packing) and the point-cloud branch of load_raw_image_array.""" +import numpy as np +import pytest + +from PIL import Image + +from weightslab.data.point_cloud_utils import ( + boxes_dimensionality, + is_point_cloud_detection_task, + box_format_string, + colorize_from_image, + compute_point_normals, + filter_valid_points, + get_pc_range, + get_point_feature_names, + is_point_cloud_task, + looks_like_point_cloud, + pack_point_cloud, + point_cloud_to_bev_image, + point_cloud_to_range_image, + point_distances, + project_boxes_to_bev, + register_boxes_fn, + register_thumbnail_fn, + render_bev_for_dataset, + render_thumbnail_2d_for_dataset, + serialize_pointcloud_box_payload, + voxel_downsample, +) + + +PC_RANGE = (0.0, -32.0, -3.0, 64.0, 32.0, 1.0) + + +def _cloud(n=1000, seed=0): + rng = np.random.default_rng(seed) + pts = np.stack([ + rng.uniform(0, 64, n), + rng.uniform(-32, 32, n), + rng.uniform(-2.0, 0.5, n), + rng.uniform(0, 1, n), + ], axis=1).astype(np.float32) + return pts + + +class _FakeDataset: + task_type = "detection_pointcloud" + pc_range = PC_RANGE + + +# --------------------------------------------------------------------------- +# Heuristics +# --------------------------------------------------------------------------- +def test_is_point_cloud_task(): + assert is_point_cloud_task("detection_pointcloud") + assert is_point_cloud_task("detection_3d") + assert is_point_cloud_task("Detection_3D") + assert is_point_cloud_task("pointcloud_seg") + assert not is_point_cloud_task("detection") + assert not is_point_cloud_task("segmentation") + assert not is_point_cloud_task(None) + + +def test_is_point_cloud_detection_task(): + assert is_point_cloud_detection_task("detection_pointcloud") + assert is_point_cloud_detection_task("Detection_PointCloud") + assert is_point_cloud_detection_task("detection_3d") # legacy alias + assert not is_point_cloud_detection_task("detection") + assert not is_point_cloud_detection_task("segmentation") + assert not is_point_cloud_detection_task(None) + + +def test_looks_like_point_cloud(): + assert looks_like_point_cloud(_cloud()) + assert looks_like_point_cloud(_cloud()[:, :3]) + assert looks_like_point_cloud(_cloud()[:, :2]) + # Multi-channel clouds (xyz + intensity + normals + rgb = 10 cols) qualify. + assert looks_like_point_cloud(np.zeros((100, 10), np.float32)) + assert not looks_like_point_cloud(_cloud()[:8]) # too few rows + assert not looks_like_point_cloud(np.zeros((100, 20), np.float32)) # too many cols + assert not looks_like_point_cloud(np.zeros((64, 64), np.uint8)) # int image + assert not looks_like_point_cloud(np.zeros((64, 64, 3), np.float32)) # 3D array + + +def test_point_distances(): + pts = np.array([[3.0, 4.0, 0.0, 1.0], [0.0, 0.0, 5.0, 0.5]], np.float32) + d = point_distances(pts) + np.testing.assert_allclose(d, [5.0, 5.0], rtol=1e-5) + + +def test_compute_point_normals_planar(): + # Points on the z=0 plane -> normals should be ~+/-z. + rng = np.random.default_rng(0) + xy = rng.uniform(-5, 5, (500, 2)).astype(np.float32) + pts = np.concatenate([xy, np.zeros((500, 1), np.float32)], axis=1) + normals = compute_point_normals(pts, k=12) + assert normals.shape == (500, 3) + np.testing.assert_allclose(np.linalg.norm(normals, axis=1), 1.0, atol=1e-4) + assert np.abs(normals[:, 2]).mean() > 0.95 # mostly aligned with z + + +def test_voxel_downsample_reduces_points(): + rng = np.random.default_rng(1) + pts = rng.uniform(0, 1, (5000, 4)).astype(np.float32) + out = voxel_downsample(pts, voxel_size=0.25) + assert out.shape[1] == 4 + assert out.shape[0] < pts.shape[0] + assert out.shape[0] <= 4 ** 3 # at most one point per 0.25 voxel in the unit cube + + +def test_colorize_from_image(): + image = np.zeros((10, 20, 3), np.uint8) + image[:, :, 0] = 255 # all red + pts = np.array([[1.0, 0.0, 0.0], [2.0, 0.0, 0.0]], np.float32) + + def project(p): + uv = np.stack([np.full(len(p), 5.0), np.full(len(p), 5.0)], axis=1) + return uv, np.array([True, False]) + + rgb = colorize_from_image(pts, image, project) + np.testing.assert_allclose(rgb[0], [1.0, 0.0, 0.0], atol=1e-5) # sampled red + np.testing.assert_allclose(rgb[1], [0.5, 0.5, 0.5], atol=1e-5) # invalid -> grey + + +def test_range_image_shape(): + img = point_cloud_to_range_image(_cloud(2000), image_height=48, image_width=256) + assert img.size == (256, 48) + arr = np.asarray(img) + assert (arr != arr[0, 0]).any() # some points were projected + + +def test_get_point_feature_names_from_dataset_and_default(): + class DS: + point_feature_names = ["x", "y", "z", "intensity", "nx", "ny", "nz"] + assert get_point_feature_names(DS(), 7) == ["x", "y", "z", "intensity", "nx", "ny", "nz"] + # Defaults when the dataset declares none. + assert get_point_feature_names(object(), 4) == ["x", "y", "z", "intensity"] + assert get_point_feature_names(object(), 3) == ["x", "y", "z"] + + +def test_registered_thumbnail_and_boxes_fns(): + marker = {"called": False} + + def my_thumb(points): + marker["called"] = True + return np.full((16, 16, 3), 9, np.uint8) + + register_thumbnail_fn(my_thumb) + try: + img = render_thumbnail_2d_for_dataset(object(), _cloud()) + assert marker["called"] and np.asarray(img)[0, 0, 0] == 9 + finally: + register_thumbnail_fn(None) # reset global state + + def my_boxes(boxes): + return np.zeros((len(boxes), 6), np.float32) + + register_boxes_fn(my_boxes) + try: + from weightslab.data.point_cloud_utils import project_boxes_for_dataset + out = project_boxes_for_dataset(object(), np.ones((3, 9), np.float32)) + assert out.shape == (3, 6) + finally: + register_boxes_fn(None) + + +def test_filter_valid_points_drops_pads_and_nonfinite(): + pts = _cloud(100) + pts[10] = -1000.0 # pad row (all coords at PAD_VALUE) + pts[20, 2] = np.nan + out = filter_valid_points(pts) + assert out.shape[0] == 98 + + +# --------------------------------------------------------------------------- +# BEV image +# --------------------------------------------------------------------------- +def test_point_cloud_to_bev_image_shape_and_content(): + img = point_cloud_to_bev_image(_cloud(), pc_range=PC_RANGE, image_size=128) + assert isinstance(img, Image.Image) + assert img.size == (128, 128) + arr = np.asarray(img) + # Some pixels must differ from the background (points were splatted). + assert (arr != arr[0, 0]).any() + + +def test_bev_image_empty_cloud_is_background_only(): + img = point_cloud_to_bev_image(np.zeros((0, 4), np.float32), pc_range=PC_RANGE, image_size=64) + arr = np.asarray(img) + assert (arr == arr[0, 0]).all() + + +def test_render_bev_for_dataset_honors_hook(): + class HookedDataset(_FakeDataset): + def to_bev_image(self, points): + return np.full((32, 32, 3), 7, np.uint8) + + img = render_bev_for_dataset(HookedDataset(), _cloud()) + assert img.size == (32, 32) + assert np.asarray(img)[0, 0, 0] == 7 + + +# --------------------------------------------------------------------------- +# Box projection +# --------------------------------------------------------------------------- +def test_project_boxes_to_bev_3d_geometry(): + # Axis-aligned box centered mid-range: easy to check normalized coords. + boxes = np.array([[32.0, 0.0, -1.0, 4.0, 2.0, 1.5, 0.0, 1.0, 0.9]], np.float32) + bev = project_boxes_to_bev(boxes, PC_RANGE, min_norm_size=0.0) + assert bev.shape == (1, 6) + x1, y1, x2, y2, cls, conf = bev[0] + assert x1 == pytest.approx((32 - 2 - 0) / 64.0, abs=1e-5) + assert x2 == pytest.approx((32 + 2 - 0) / 64.0, abs=1e-5) + # y axis flips (image v grows downward): cy=0 -> centered. + assert (y1 + y2) / 2 == pytest.approx(0.5, abs=1e-5) + assert cls == 1.0 and conf == pytest.approx(0.9) + + +def test_project_boxes_yaw_rotation_grows_extent(): + no_yaw = project_boxes_to_bev( + np.array([[32, 0, -1, 4.0, 2.0, 1.5, 0.0, 0, 1]], np.float32), PC_RANGE, 0.0) + yawed = project_boxes_to_bev( + np.array([[32, 0, -1, 4.0, 2.0, 1.5, np.pi / 4, 0, 1]], np.float32), PC_RANGE, 0.0) + assert (yawed[0, 2] - yawed[0, 0]) > (no_yaw[0, 2] - no_yaw[0, 0]) - 1e-6 + assert (yawed[0, 3] - yawed[0, 1]) > (no_yaw[0, 3] - no_yaw[0, 1]) + + +def test_project_boxes_min_size_clamp(): + tiny = np.array([[32.0, 0.0, -1.0, 0.01, 0.01, 0.01, 0.0, 0, 1]], np.float32) + bev = project_boxes_to_bev(tiny, PC_RANGE, min_norm_size=0.01) + assert (bev[0, 2] - bev[0, 0]) >= 0.0099 + assert (bev[0, 3] - bev[0, 1]) >= 0.0099 + + +def test_project_boxes_2d_rows(): + boxes = np.array([[10.0, 5.0, 2.0, 2.0, 2.0, 0.7]], np.float32) # cx,cy,dx,dy,cls,conf + assert boxes_dimensionality(boxes) == 2 + bev = project_boxes_to_bev(boxes, PC_RANGE, 0.0) + assert bev[0, 4] == 2.0 + assert bev[0, 5] == pytest.approx(0.7) + + +def test_box_format_string(): + assert box_format_string(np.zeros((1, 9), np.float32)) == "cx_cy_cz_dx_dy_dz_yaw_cls_conf" + assert box_format_string(np.zeros((1, 6), np.float32)) == "cx_cy_dx_dy_cls_conf" + + +# --------------------------------------------------------------------------- +# Payload + range resolution +# --------------------------------------------------------------------------- +def test_serialize_pointcloud_box_payload(): + ds = _FakeDataset() + boxes = np.array([ + [32.0, 0.0, -1.0, 4.0, 2.0, 1.5, 0.3, 1.0, 0.8], + [10.0, -5.0, -1.2, 0.8, 0.6, 1.7, -1.0, 2.0, 0.5], + ], np.float32) + payload = serialize_pointcloud_box_payload(ds, boxes) + assert payload["format"] == "xyxy" + assert len(payload["bboxes"]) == 2 and len(payload["bboxes"][0]) == 6 + assert len(payload["bboxes_3d"]) == 2 and len(payload["bboxes_3d"][0]) == 9 + assert payload["pc_range"] == list(PC_RANGE) + + +def test_get_pc_range_attr_and_auto(): + assert get_pc_range(_FakeDataset()) == PC_RANGE + + class Bare: + pass + bare = Bare() + auto = get_pc_range(bare, _cloud()) + assert auto is not None and len(auto) == 6 + # Cached on the dataset for later (image/box alignment). + assert get_pc_range(bare) == auto + + +# --------------------------------------------------------------------------- +# Binary packing (GetPointCloud) +# --------------------------------------------------------------------------- +def test_pack_point_cloud_roundtrip_and_downsample(): + pts = _cloud(5000) + data, n, f = pack_point_cloud(pts, max_points=0) + assert (n, f) == (5000, 4) + decoded = np.frombuffer(data, dtype="= 1) +# --------------------------------------------------------------------------- +class TestSaveInstanceSignals(_SaveSignalsBase): + def test_segmentation_flat_signals_and_nested_targets(self): + """Dense (segmentation) per-instance: flat signals + nested-list mask targets + route to (sid, 1..N) with distinct values per instance.""" + mgr = _fresh_manager(["0", "1"]) + masks = [np.full((8, 8), i + 1, dtype=np.uint8) for i in range(3)] + v = self._call( + src.save_instance_signals, mgr, + signals={"iou": th.tensor([0.71, 0.82, 0.93])}, + batch_ids=["0", "1"], + batch_idx=th.tensor([0, 0, 1]), # sample 0: 2 instances, sample 1: 1 + targets=[[masks[0], masks[1]], [masks[2]]], + origin="train", + log=False, + ) + # Sample rows + 1-based instance rows. + self.assertEqual(sorted(v.loc["0"].index.tolist()), [0, 1, 2]) + self.assertEqual(sorted(v.loc["1"].index.tolist()), [0, 1]) + # Signals aligned per instance (flat, sample-major). + self.assertAlmostEqual(float(v.loc[("0", 1), "signals//iou"]), 0.71, places=5) + self.assertAlmostEqual(float(v.loc[("0", 2), "signals//iou"]), 0.82, places=5) + self.assertAlmostEqual(float(v.loc[("1", 1), "signals//iou"]), 0.93, places=5) + # Each instance carries its own mask target. + self.assertEqual(int(np.median(np.asarray(v.loc[("0", 1), TARGET]))), 1) + self.assertEqual(int(np.median(np.asarray(v.loc[("0", 2), TARGET]))), 2) + self.assertEqual(int(np.median(np.asarray(v.loc[("1", 1), TARGET]))), 3) + + def test_detection_dict_targets(self): + """Dict (Ultralytics) per-instance: bbox+cls targets are split per instance, + and flat signals route to the matching (sid, annotation_id).""" + mgr = _fresh_manager(["0", "1"]) + tdict = { + "batch_idx": th.tensor([0, 0, 1]), + "bboxes": th.tensor([[1, 1, 3, 3], [5, 5, 7, 7], [2, 2, 4, 4]], dtype=th.float32), + "cls": th.tensor([[1.0], [2.0], [3.0]]), + } + v = self._call( + src.save_instance_signals, mgr, + signals={"iou": th.tensor([0.71, 0.82, 0.93])}, + batch_ids=["0", "1"], + batch_idx=th.tensor([0, 0, 1]), + targets=tdict, + origin="train", + log=False, + ) + self.assertEqual(sorted(v.loc["0"].index.tolist()), [0, 1, 2]) + self.assertEqual(sorted(v.loc["1"].index.tolist()), [0, 1]) + # Per-instance signals correct (not mis-routed to a single instance). + self.assertAlmostEqual(float(v.loc[("0", 1), "signals//iou"]), 0.71, places=5) + self.assertAlmostEqual(float(v.loc[("0", 2), "signals//iou"]), 0.82, places=5) + self.assertAlmostEqual(float(v.loc[("1", 1), "signals//iou"]), 0.93, places=5) + # Each instance target = its box coords + class id. + np.testing.assert_array_equal(np.asarray(v.loc[("0", 1), TARGET]), [1, 1, 3, 3, 1]) + np.testing.assert_array_equal(np.asarray(v.loc[("0", 2), TARGET]), [5, 5, 7, 7, 2]) + np.testing.assert_array_equal(np.asarray(v.loc[("1", 1), TARGET]), [2, 2, 4, 4, 3]) + + def test_instance_signals_do_not_touch_sample_row(self): + """The per-sample row (annotation_id 0) keeps NaN for a per-instance signal.""" + mgr = _fresh_manager(["0"]) + v = self._call( + src.save_instance_signals, mgr, + signals={"iou": th.tensor([0.5, 0.6])}, + batch_ids=["0"], + batch_idx=th.tensor([0, 0]), + origin="train", + log=False, + ) + self.assertTrue(np.isnan(float(v.loc[("0", 0), "signals//iou"]))) + self.assertAlmostEqual(float(v.loc[("0", 1), "signals//iou"]), 0.5, places=5) + self.assertAlmostEqual(float(v.loc[("0", 2), "signals//iou"]), 0.6, places=5) + + def test_empty_batch_idx_is_a_noop(self): + """No instances → nothing enqueued, no crash.""" + mgr = _fresh_manager(["0"]) + v = self._call( + src.save_instance_signals, mgr, + signals={"iou": th.tensor([])}, + batch_ids=["0"], + batch_idx=th.tensor([], dtype=th.long), + origin="train", + log=False, + ) + self.assertEqual(sorted(v.loc["0"].index.tolist()), [0]) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/trainer/services/test_notebook_service_unit.py b/tests/trainer/services/test_notebook_service_unit.py index b8a6adfa..e8c73825 100644 --- a/tests/trainer/services/test_notebook_service_unit.py +++ b/tests/trainer/services/test_notebook_service_unit.py @@ -395,6 +395,54 @@ def _make_service(self): return NotebookService(_fake_data_service(), root_log_dir=str(self.root)) +class TestEmbeddedKernelStartupOrder(unittest.TestCase): + """Ordering invariants in ``_run_embedded_kernel``'s startup sequence. + + Asserted against the source rather than a live kernel on purpose: the + things being ordered are process-wide ``sys.stdout``/``sys.stderr`` swaps, + and pytest's own capture replaces those again per test, so inspecting them + from inside a test reports whatever capture did last, not what the kernel + did. Reading the order directly is deterministic and says exactly what the + constraint is. + """ + + def _source(self): + import inspect + return inspect.getsource(notebook_service._run_embedded_kernel) + + def test_flush_interval_is_lowered_before_the_streams_are_wrapped(self): + # ipykernel's OutStream batches writes to iopub every `flush_interval` + # seconds (0.2s default); the kernel lowers it to 0.05s so cell output + # streams live instead of arriving in one lump at the end. + # + # It has to be set while sys.stdout is still the raw OutStream. + # _ThreadRoutedStream.__getattr__ delegates reads to the stream it + # wraps, so `hasattr(..., "flush_interval")` keeps answering True after + # wrapping and the loop still "works" -- but the assignment lands on + # the wrapper and the real OutStream keeps its 0.2s. The live streaming + # test catches that only through wall-clock timing, which makes it a + # close call rather than a clear failure. + source = self._source() + flush_at = source.index("flush_interval = 0.05") + wrap_at = source.index("_install_thread_routed_streams(") + self.assertLess( + flush_at, wrap_at, + "flush_interval must be set on the raw OutStreams, before " + "_install_thread_routed_streams wraps them") + + def test_logging_is_repaired_before_anything_else_logs(self): + # IPKernelApp.initialize() -> traitlets -> logging.config.dictConfig + # closes every handler in the process, killing the session log file. + source = self._source() + self.assertLess( + source.index("app.initialize("), source.index("ensure_logging_intact()"), + "logging can only be repaired after initialize() has broken it") + self.assertLess( + source.index("ensure_logging_intact()"), + source.index("_install_thread_routed_streams("), + "repair logging before the rest of the startup logs anything") + + class TestNotebookPersistence(unittest.TestCase): def setUp(self): self._tmp = tempfile.TemporaryDirectory() diff --git a/tests/utils/test_logs_unit.py b/tests/utils/test_logs_unit.py index 2c11204b..3218024f 100644 --- a/tests/utils/test_logs_unit.py +++ b/tests/utils/test_logs_unit.py @@ -1,3 +1,4 @@ +import io import logging import os import shutil @@ -8,8 +9,29 @@ from weightslab.utils import logs -class TestLogsUnit(unittest.TestCase): +class LogsTestBase(unittest.TestCase): + """Isolates the root logger and the module globals from the test runner's.""" + def setUp(self): + self._tmpdirs = [] + self._saved_env = { + key: os.environ.pop(key, None) + for key in ("WEIGHTSLAB_ROOT_LOG_DIR", "WEIGHTSLAB_LOG_FILE_LEVEL", + "WEIGHTSLAB_TQDM_LOG_TO_TERMINAL") + } + self._reset_logging() + + def tearDown(self): + self._reset_logging() + for key, value in self._saved_env.items(): + if value is None: + os.environ.pop(key, None) + else: + os.environ[key] = value + for path in self._tmpdirs: + shutil.rmtree(path, ignore_errors=True) + + def _reset_logging(self): root = logging.getLogger() for handler in list(root.handlers): try: @@ -21,44 +43,45 @@ def setUp(self): logs._LOG_FILE_PATH = None logs._TMP_DIR_PATH = None logs._FILE_HANDLER = None + logs._CONSOLE_HANDLER = None + logs._CONSOLE_LEVEL = logging.INFO + logs._EXIT_HOOK_REGISTERED = False + logs._FILE_LEVEL = logging.NOTSET + + def _mkdtemp(self): + path = tempfile.mkdtemp() + self._tmpdirs.append(path) + return path + + def _log_file_contents(self): + logs.flush_logs() + with open(logs._LOG_FILE_PATH, encoding="utf-8") as handle: + return handle.read() - def tearDown(self): - root = logging.getLogger() - for handler in list(root.handlers): - try: - handler.close() - except Exception: - pass - root.removeHandler(handler) +class TestLogsUnit(LogsTestBase): def test_setup_logging_with_file_and_print_location(self): logs.setup_logging("INFO", log_to_file=True) self.assertIsNotNone(logs._LOG_FILE_PATH) self.assertTrue(os.path.exists(logs._LOG_FILE_PATH)) - with patch("weightslab.utils.logs.print") as p: + # builtins.print, not the module's logging shim: the notice has to reach + # the terminal even once logging is being torn down at interpreter exit. + with patch("weightslab.utils.logs.builtins.print") as p: logs._print_log_location() p.assert_called_once() + self.assertIn(logs._LOG_FILE_PATH, p.call_args.args[0]) def test_set_log_directory_moves_log_and_reopens_handler(self): logs.setup_logging("DEBUG", log_to_file=True) old_path = logs._LOG_FILE_PATH - tmpdir = tempfile.mkdtemp() - try: - logs.set_log_directory(tmpdir) - self.assertNotEqual(old_path, logs._LOG_FILE_PATH) - self.assertTrue(logs._LOG_FILE_PATH.startswith(tmpdir)) - self.assertTrue(os.path.exists(logs._LOG_FILE_PATH)) - finally: - if logs._FILE_HANDLER is not None: - try: - logs._FILE_HANDLER.flush() - logs._FILE_HANDLER.close() - except Exception: - pass - logging.getLogger().handlers = [] - shutil.rmtree(tmpdir, ignore_errors=True) + tmpdir = self._mkdtemp() + logs.set_log_directory(tmpdir) + self.assertNotEqual(old_path, logs._LOG_FILE_PATH) + self.assertTrue(logs._LOG_FILE_PATH.startswith(tmpdir)) + self.assertTrue(os.path.exists(logs._LOG_FILE_PATH)) + self.assertFalse(os.path.exists(old_path)) def test_custom_print_routes_to_levels(self): with patch("logging.info") as info_mock, \ @@ -75,5 +98,285 @@ def test_set_log_directory_without_setup_is_noop(self): warn_mock.assert_called_once() +class TestLevelResolution(LogsTestBase): + def test_resolve_level_accepts_names_numbers_and_garbage(self): + self.assertEqual(logs._resolve_level("debug"), logging.DEBUG) + self.assertEqual(logs._resolve_level("WARNING"), logging.WARNING) + self.assertEqual(logs._resolve_level(" Error "), logging.ERROR) + self.assertEqual(logs._resolve_level(25), 25) + self.assertEqual(logs._resolve_level("25"), 25) + # Unknown / empty / None fall back instead of raising, so a typo in + # WEIGHTSLAB_LOG_LEVEL degrades logging rather than killing the import. + self.assertEqual(logs._resolve_level("nonsense"), logging.INFO) + self.assertEqual(logs._resolve_level(""), logging.INFO) + self.assertEqual(logs._resolve_level(None), logging.INFO) + self.assertEqual(logs._resolve_level(None, default=logging.NOTSET), logging.NOTSET) + + def test_resolve_level_knows_the_custom_watchdog_level(self): + from weightslab.watchdog.log_level import WATCHDOG + self.assertEqual(logs._resolve_level("WATCHDOG"), WATCHDOG) + + +class TestConsoleAndFileLevelsAreIndependent(LogsTestBase): + """The terminal honours WEIGHTSLAB_LOG_LEVEL; the file keeps everything.""" + + def test_file_keeps_debug_when_console_is_info(self): + logs.setup_logging("INFO", log_to_file=True) + + # The root logger gates records before any handler sees them, so it has + # to be wide open or the file handler's own level is moot. + self.assertEqual(logging.getLogger().level, logging.NOTSET) + self.assertEqual(logs._CONSOLE_HANDLER.level, logging.INFO) + self.assertEqual(logs._FILE_HANDLER.level, logging.NOTSET) + + logging.getLogger("test.independent").debug("debug-marker") + logging.getLogger("test.independent").info("info-marker") + + contents = self._log_file_contents() + self.assertIn("debug-marker", contents) + self.assertIn("info-marker", contents) + + def test_terminal_filters_what_the_file_still_records(self): + logs.setup_logging("WARNING", log_to_file=True) + terminal = io.StringIO() + logs._CONSOLE_HANDLER.setStream(terminal) + + log = logging.getLogger("test.console") + log.debug("debug-marker") + log.info("info-marker") + log.warning("warning-marker") + + printed = terminal.getvalue() + self.assertNotIn("debug-marker", printed) + self.assertNotIn("info-marker", printed) + self.assertIn("warning-marker", printed) + + # All three are on disk regardless of what the terminal showed. + contents = self._log_file_contents() + for marker in ("debug-marker", "info-marker", "warning-marker"): + self.assertIn(marker, contents) + + def test_file_level_can_be_capped_by_env_var(self): + os.environ["WEIGHTSLAB_LOG_FILE_LEVEL"] = "WARNING" + logs.setup_logging("DEBUG", log_to_file=True) + + self.assertEqual(logs._FILE_HANDLER.level, logging.WARNING) + # Root sits at the most permissive of the two sinks (DEBUG here). + self.assertEqual(logging.getLogger().level, logging.DEBUG) + + logging.getLogger("test.capped").debug("debug-marker") + logging.getLogger("test.capped").warning("warning-marker") + + contents = self._log_file_contents() + self.assertNotIn("debug-marker", contents) + self.assertIn("warning-marker", contents) + + def test_file_level_argument_overrides_env_var(self): + os.environ["WEIGHTSLAB_LOG_FILE_LEVEL"] = "WARNING" + logs.setup_logging("INFO", log_to_file=True, file_level="DEBUG") + self.assertEqual(logs._FILE_HANDLER.level, logging.DEBUG) + + def test_without_file_logging_root_keeps_the_console_level(self): + logs.setup_logging("INFO", log_to_file=False) + self.assertEqual(logging.getLogger().level, logging.INFO) + self.assertIsNone(logs._FILE_HANDLER) + + +class TestLogDirectoryLayout(LogsTestBase): + """A session log must never move between two directory conventions.""" + + def test_experiment_log_dir_is_the_root_plus_subdir(self): + self.assertEqual( + logs.experiment_log_dir(os.path.join("a", "b")), + os.path.join("a", "b", logs.LOG_SUBDIR)) + + def test_setup_logging_uses_the_root_log_dir_env_var(self): + root = self._mkdtemp() + os.environ["WEIGHTSLAB_ROOT_LOG_DIR"] = root + logs.setup_logging("INFO", log_to_file=True) + + self.assertEqual( + os.path.dirname(logs._LOG_FILE_PATH), logs.experiment_log_dir(root)) + + def test_relocation_lands_in_the_same_subdir_setup_logging_uses(self): + logs.setup_logging("INFO", log_to_file=True) + root = self._mkdtemp() + + logs.set_log_directory(logs.experiment_log_dir(root)) + + self.assertEqual( + os.path.dirname(logs._LOG_FILE_PATH), logs.experiment_log_dir(root)) + self.assertTrue(os.path.exists(logs._LOG_FILE_PATH)) + + def test_relocation_carries_history_over_and_keeps_appending(self): + logs.setup_logging("INFO", log_to_file=True) + logging.getLogger("test.move").info("before-marker") + + logs.set_log_directory(self._mkdtemp()) + logging.getLogger("test.move").info("after-marker") + + contents = self._log_file_contents() + self.assertIn("before-marker", contents) + self.assertIn("after-marker", contents) + + def test_relocation_preserves_the_file_level(self): + logs.setup_logging("INFO", log_to_file=True) + logs.set_log_directory(self._mkdtemp()) + + self.assertEqual(logs._FILE_HANDLER.level, logging.NOTSET) + logging.getLogger("test.level").debug("debug-after-move") + self.assertIn("debug-after-move", self._log_file_contents()) + + def test_relocating_to_the_current_directory_is_a_noop(self): + logs.setup_logging("INFO", log_to_file=True) + path = logs._LOG_FILE_PATH + handler = logs._FILE_HANDLER + + logs.set_log_directory(os.path.dirname(path)) + + # Same file, same handler: no churn and no repeated "updated" lines. + self.assertEqual(logs._LOG_FILE_PATH, path) + self.assertIs(logs._FILE_HANDLER, handler) + + def test_failed_move_still_leaves_a_working_handler(self): + logs.setup_logging("INFO", log_to_file=True) + target = self._mkdtemp() + + with patch("weightslab.utils.logs.shutil.move", side_effect=OSError("locked")): + logs.set_log_directory(target) + + self.assertEqual(os.path.dirname(logs._LOG_FILE_PATH), target) + logging.getLogger("test.failed_move").info("still-logging") + self.assertIn("still-logging", self._log_file_contents()) + + +class TestSurvivesExternalLoggingReconfiguration(LogsTestBase): + """``logging.config.dictConfig`` closes every handler in the process. + + traitlets runs it while building any ``Application`` — which is what + ipykernel does every time the studio's embedded notebook kernel starts. + ``_clearExistingHandlers`` closes the handlers but leaves them attached to + the root logger, so a dead one goes on accepting records and dropping them. + """ + + @staticmethod + def _external_dictconfig(): + import logging.config + logging.config.dictConfig({ + "version": 1, "handlers": {}, "loggers": {}, + "disable_existing_loggers": False, + }) + + def test_session_log_keeps_recording_without_any_repair_call(self): + logs.setup_logging("INFO", log_to_file=True) + logging.getLogger("test.reconfig").info("before-marker") + + self._external_dictconfig() + logging.getLogger("test.reconfig").info("after-marker") + + contents = self._log_file_contents() + self.assertIn("before-marker", contents) + self.assertIn("after-marker", contents) + + def test_a_mode_w_handler_would_have_lost_the_record(self): + # Guards the reason _make_file_handler appends: CPython refuses to + # reopen a closed mode='w' FileHandler (bpo-42378), so the old handler + # died silently right here. + logs.setup_logging("INFO", log_to_file=True) + root = logging.getLogger() + root.removeHandler(logs._FILE_HANDLER) + logs._FILE_HANDLER.close() + truncating = logging.FileHandler(logs._LOG_FILE_PATH, mode="w", encoding="utf-8") + truncating.setFormatter(logging.Formatter(logs.FORMAT, datefmt=logs.DATE_FORMAT)) + root.addHandler(truncating) + logs._FILE_HANDLER = truncating + + self._external_dictconfig() + logging.getLogger("test.reconfig").info("after-marker") + + self.assertIn(truncating, root.handlers) # still attached... + self.assertIsNone(truncating.stream) # ...but stream gone + self.assertNotIn("after-marker", self._log_file_contents()) + + def test_ensure_logging_intact_reopens_a_closed_file_handler(self): + logs.setup_logging("INFO", log_to_file=True) + logs._FILE_HANDLER.close() + self.assertIsNone(logs._FILE_HANDLER.stream) + + self.assertTrue(logs.ensure_logging_intact()) + + logging.getLogger("test.repair").info("repaired-marker") + self.assertIn("repaired-marker", self._log_file_contents()) + + def test_ensure_logging_intact_reattaches_detached_handlers(self): + logs.setup_logging("INFO", log_to_file=True) + root = logging.getLogger() + root.handlers = [] + + self.assertTrue(logs.ensure_logging_intact()) + + self.assertIn(logs._CONSOLE_HANDLER, root.handlers) + self.assertIn(logs._FILE_HANDLER, root.handlers) + + def test_ensure_logging_intact_restores_the_root_level(self): + logs.setup_logging("INFO", log_to_file=True) + logging.getLogger().setLevel(logging.CRITICAL) + + self.assertTrue(logs.ensure_logging_intact()) + self.assertEqual(logging.getLogger().level, logging.NOTSET) + + def test_ensure_logging_intact_is_a_noop_when_nothing_is_broken(self): + logs.setup_logging("INFO", log_to_file=True) + self.assertFalse(logs.ensure_logging_intact()) + + def test_ensure_logging_intact_without_setup_does_not_raise(self): + self.assertFalse(logs.ensure_logging_intact()) + + +class TestProgressChannel(LogsTestBase): + """The tqdm mirror belongs in the file, not on top of the live bar.""" + + def test_progress_records_reach_the_file_but_not_the_terminal(self): + logs.setup_logging("INFO", log_to_file=True) + terminal = io.StringIO() + logs._CONSOLE_HANDLER.setStream(terminal) + + logging.getLogger(logs.PROGRESS_LOGGER_NAME).info("Training: 120 steps | loss=0.4") + logging.getLogger("test.ordinary").info("ordinary-marker") + + printed = terminal.getvalue() + self.assertNotIn("Training: 120 steps", printed) + self.assertIn("ordinary-marker", printed) + + contents = self._log_file_contents() + self.assertIn("Training: 120 steps", contents) + + def test_the_terminal_handler_never_prints_progress_itself(self): + # The terminal copy comes from tqdm.write (see tqdm_logging), which + # redraws the bar around it. A handler write would land mid-bar, so the + # filter stays on regardless of WEIGHTSLAB_TQDM_LOG_TO_TERMINAL. + os.environ["WEIGHTSLAB_TQDM_LOG_TO_TERMINAL"] = "1" + logs.setup_logging("INFO", log_to_file=True) + terminal = io.StringIO() + logs._CONSOLE_HANDLER.setStream(terminal) + + logging.getLogger(logs.PROGRESS_LOGGER_NAME).info("Training: 120 steps") + self.assertNotIn("Training: 120 steps", terminal.getvalue()) + self.assertIn("Training: 120 steps", self._log_file_contents()) + + +class TestFlushLogs(LogsTestBase): + def test_flush_logs_without_a_handler_is_a_noop(self): + logs.flush_logs() # must not raise + + def test_flush_logs_pushes_records_to_disk(self): + logs.setup_logging("INFO", log_to_file=True) + logging.getLogger("test.flush").info("flushed-marker") + logs.flush_logs() + + with open(logs._LOG_FILE_PATH, encoding="utf-8") as handle: + self.assertIn("flushed-marker", handle.read()) + + if __name__ == "__main__": unittest.main() diff --git a/tests/utils/test_tqdm_logging_unit.py b/tests/utils/test_tqdm_logging_unit.py new file mode 100644 index 00000000..487573bb --- /dev/null +++ b/tests/utils/test_tqdm_logging_unit.py @@ -0,0 +1,231 @@ +"""Tests for the tqdm -> session log mirror. + +A tqdm bar paints itself onto a terminal and never touches ``logging``, so the +log file had no record of a run's own progress. These cover the sampler that +puts it there. +""" + +import logging +import os +import unittest +from unittest.mock import patch + +from weightslab.utils import tqdm_logging +from weightslab.utils.logs import PROGRESS_LOGGER_NAME + + +class _FakeBar: + """Stands in for a tqdm instance: only ``format_dict`` is read.""" + + def __init__(self, **fields): + base = {"prefix": "Training", "n": 0, "total": None, + "elapsed": 0.0, "rate": None, "postfix": None} + base.update(fields) + self.format_dict = base + + +_ENV_KEYS = ("WEIGHTSLAB_TQDM_LOG_INTERVAL", "WEIGHTSLAB_TQDM_LOG_TO_TERMINAL") + + +class TqdmLoggingTestBase(unittest.TestCase): + def setUp(self): + self._saved = {key: os.environ.pop(key, None) for key in _ENV_KEYS} + # Off by default in tests: the echo is verified explicitly below, and + # every other test would otherwise print into the runner's output. + os.environ["WEIGHTSLAB_TQDM_LOG_TO_TERMINAL"] = "0" + self.addCleanup(self._restore_env) + self.addCleanup(tqdm_logging.stop_tqdm_log_mirror) + + def _restore_env(self): + for key, value in self._saved.items(): + if value is None: + os.environ.pop(key, None) + else: + os.environ[key] = value + + +class TestRender(TqdmLoggingTestBase): + def test_renders_a_compact_line_without_the_drawn_bar(self): + bar = _FakeBar(n=2750, elapsed=1752.0, rate=1.57, + postfix="train_loss=0.3421 | test_acc=88.3%") + line = tqdm_logging._render(bar) + + self.assertIn("Training", line) + self.assertIn("2750 steps", line) + self.assertIn("1.57 it/s", line) + self.assertIn("train_loss=0.3421", line) + # The block-drawing bar carries nothing in a log file. + self.assertNotIn("█", line) + self.assertNotIn("\r", line) + + def test_shows_a_percentage_when_the_total_is_known(self): + line = tqdm_logging._render(_FakeBar(n=25, total=100, elapsed=10.0)) + self.assertIn("25/100", line) + self.assertIn("25%", line) + + def test_falls_back_to_a_generic_label_without_a_description(self): + self.assertTrue(tqdm_logging._render(_FakeBar(prefix=None)).startswith("progress:")) + + def test_a_bar_that_cannot_be_read_is_skipped_not_raised(self): + class Broken: + @property + def format_dict(self): + raise RuntimeError("gone") + + self.assertIsNone(tqdm_logging._render(Broken())) + + +class TestSampling(TqdmLoggingTestBase): + def test_logs_one_line_per_bar_to_the_progress_channel(self): + bar = _FakeBar(n=10, elapsed=5.0, postfix="train_loss=0.5") + with patch.object(tqdm_logging, "_live_bars", return_value=[bar]), \ + self.assertLogs(PROGRESS_LOGGER_NAME, level=logging.INFO) as captured: + tqdm_logging._sample_once({}) + + self.assertEqual(len(captured.records), 1) + self.assertIn("train_loss=0.5", captured.output[0]) + + def test_an_unchanged_bar_is_not_logged_twice(self): + bar = _FakeBar(n=10, elapsed=5.0) + seen = {} + with patch.object(tqdm_logging, "_live_bars", return_value=[bar]): + with self.assertLogs(PROGRESS_LOGGER_NAME, level=logging.INFO): + tqdm_logging._sample_once(seen) + # A paused run must not fill the log with identical lines. + with patch.object(tqdm_logging.progress_logger, "info") as info: + tqdm_logging._sample_once(seen) + info.assert_not_called() + + def test_a_paused_bar_is_not_relogged_just_because_time_passed(self): + # Regression: the skip used to compare rendered lines, which carry + # elapsed time. A stopped run therefore wrote a near-identical entry + # every interval forever. + bar = _FakeBar(n=453, elapsed=148.0, rate=2.90, postfix="train_loss=0.1861") + seen = {} + with patch.object(tqdm_logging, "_live_bars", return_value=[bar]): + with self.assertLogs(PROGRESS_LOGGER_NAME, level=logging.INFO): + tqdm_logging._sample_once(seen) + for extra in (30.0, 60.0, 90.0): + bar.format_dict["elapsed"] = 148.0 + extra + with patch.object(tqdm_logging.progress_logger, "info") as info: + tqdm_logging._sample_once(seen) + info.assert_not_called() + + def test_a_bar_that_moved_is_logged_again(self): + bar = _FakeBar(n=10, elapsed=5.0) + seen = {} + with patch.object(tqdm_logging, "_live_bars", return_value=[bar]): + with self.assertLogs(PROGRESS_LOGGER_NAME, level=logging.INFO): + tqdm_logging._sample_once(seen) + bar.format_dict["n"] = 20 + with self.assertLogs(PROGRESS_LOGGER_NAME, level=logging.INFO) as second: + tqdm_logging._sample_once(seen) + self.assertIn("20 steps", second.output[0]) + + def test_closed_bars_are_forgotten(self): + bar = _FakeBar(n=10) + seen = {} + with patch.object(tqdm_logging, "_live_bars", return_value=[bar]), \ + self.assertLogs(PROGRESS_LOGGER_NAME, level=logging.INFO): + tqdm_logging._sample_once(seen) + self.assertEqual(len(seen), 1) + + with patch.object(tqdm_logging, "_live_bars", return_value=[]): + tqdm_logging._sample_once(seen) + self.assertEqual(seen, {}) + + def test_no_bars_logs_nothing(self): + with patch.object(tqdm_logging, "_live_bars", return_value=[]), \ + patch.object(tqdm_logging.progress_logger, "info") as info: + tqdm_logging._sample_once({}) + info.assert_not_called() + + +class TestTerminalEcho(TqdmLoggingTestBase): + """Progress also reaches the terminal, without corrupting the live bar.""" + + def test_echoes_through_tqdm_write_by_default(self): + os.environ.pop("WEIGHTSLAB_TQDM_LOG_TO_TERMINAL", None) + bar = _FakeBar(n=10, elapsed=5.0, postfix="train_loss=0.5") + import tqdm as tqdm_mod + + with patch.object(tqdm_logging, "_live_bars", return_value=[bar]), \ + patch.object(tqdm_mod.tqdm, "write") as write, \ + self.assertLogs(PROGRESS_LOGGER_NAME, level=logging.INFO): + tqdm_logging._sample_once({}) + + # tqdm.write, not a bare print: it lifts the bar, prints, and redraws. + write.assert_called_once() + self.assertIn("train_loss=0.5", write.call_args.args[0]) + + def test_echo_can_be_turned_off(self): + os.environ["WEIGHTSLAB_TQDM_LOG_TO_TERMINAL"] = "0" + bar = _FakeBar(n=10, elapsed=5.0) + import tqdm as tqdm_mod + + with patch.object(tqdm_logging, "_live_bars", return_value=[bar]), \ + patch.object(tqdm_mod.tqdm, "write") as write, \ + self.assertLogs(PROGRESS_LOGGER_NAME, level=logging.INFO): + tqdm_logging._sample_once({}) + + write.assert_not_called() + + def test_a_failing_echo_does_not_stop_the_log_line(self): + import tqdm as tqdm_mod + os.environ.pop("WEIGHTSLAB_TQDM_LOG_TO_TERMINAL", None) + + with patch.object(tqdm_logging, "_live_bars", return_value=[_FakeBar(n=1)]), \ + patch.object(tqdm_mod.tqdm, "write", side_effect=OSError("closed")), \ + self.assertLogs(PROGRESS_LOGGER_NAME, level=logging.INFO) as captured: + tqdm_logging._sample_once({}) + + self.assertEqual(len(captured.records), 1) + + +class TestLifecycle(TqdmLoggingTestBase): + def test_starts_and_is_idempotent(self): + self.assertTrue(tqdm_logging.start_tqdm_log_mirror(interval=30)) + first = tqdm_logging._thread + self.assertTrue(tqdm_logging.start_tqdm_log_mirror(interval=30)) + self.assertIs(tqdm_logging._thread, first) + + def test_a_non_positive_interval_disables_it(self): + self.assertFalse(tqdm_logging.start_tqdm_log_mirror(interval=0)) + self.assertIsNone(tqdm_logging._thread) + + def test_the_interval_comes_from_the_environment(self): + os.environ["WEIGHTSLAB_TQDM_LOG_INTERVAL"] = "0" + self.assertFalse(tqdm_logging.start_tqdm_log_mirror()) + + os.environ["WEIGHTSLAB_TQDM_LOG_INTERVAL"] = "not-a-number" + self.assertEqual(tqdm_logging._interval_from_env(), + tqdm_logging.DEFAULT_INTERVAL_SECONDS) + + def test_stop_is_safe_when_never_started(self): + tqdm_logging.stop_tqdm_log_mirror() # must not raise + + def test_stop_ends_the_thread(self): + tqdm_logging.start_tqdm_log_mirror(interval=30) + thread = tqdm_logging._thread + tqdm_logging.stop_tqdm_log_mirror() + self.assertIsNone(tqdm_logging._thread) + self.assertFalse(thread.is_alive()) + + +class TestLiveBars(TqdmLoggingTestBase): + def test_reads_real_tqdm_instances(self): + import tqdm as tqdm_mod + + with open(os.devnull, "w") as sink: + bar = tqdm_mod.tqdm(total=10, desc="RealBar", file=sink) + try: + bar.update(3) + rendered = [tqdm_logging._render(b) for b in tqdm_logging._live_bars()] + finally: + bar.close() + + self.assertTrue(any(line and "RealBar" in line for line in rendered)) + + +if __name__ == "__main__": + unittest.main() diff --git a/weightslab/AGENTS.md b/weightslab/AGENTS.md index 42851eb5..1271bf18 100644 --- a/weightslab/AGENTS.md +++ b/weightslab/AGENTS.md @@ -76,8 +76,9 @@ via the decision table in §3.9, then copy its `wl.*` calls — §3 documents th whole API surface (reactive signals, group signals, the Ultralytics mixin, etc. aren't in the `.rst` docs; the examples are the primary source). -TLS/UI deploy details: `weightslab/docs/weights_studio.rst`. TLS is opt-in: -`weightslab se` once, then `weightslab start --certs`. Windows: `se` uses the +TLS/UI deploy details: `weightslab/docs/weights_studio.rst`. TLS turns on once +certs exist: `weightslab se` once, then `weightslab start` and the backend use +them automatically (`--no-certs` forces HTTP). Windows: `se` uses the PowerShell script + Windows `openssl`; `--force-ubuntu` uses WSL bash instead. --- @@ -310,10 +311,12 @@ Authoritative reference: `weightslab/docs/configuration.rst`. High-signal ones: | Variable | Default | Why | |---|---|---| -| `WEIGHTSLAB_LOG_LEVEL` | `INFO` | `DEBUG` for detail (`WATCHDOG` level sits between WARNING/ERROR). | +| `WEIGHTSLAB_LOG_LEVEL` | `INFO` | **Terminal only**; `DEBUG` for detail (`WATCHDOG` level sits between WARNING/ERROR). | +| `WEIGHTSLAB_LOG_FILE_LEVEL` | *(unset = all)* | The session log file (`/weightslab_logs/`) keeps every record whatever the terminal shows; set this to cap the file too. | +| `WEIGHTSLAB_TQDM_LOG_INTERVAL` | `30` | Seconds between snapshots of live `tqdm` bars into the log (`0` disables); tqdm never goes through `logging`. | | `GRPC_BACKEND_HOST`/`PORT` | `0.0.0.0`/`50051` | Backend gRPC bind address. | -| `GRPC_TLS_ENABLED` | `0` | TLS on the gRPC socket; set with `weightslab start --certs`. | -| `GRPC_TLS_REQUIRE_CLIENT_AUTH` | `0` | mTLS; must match `--certs`. | +| `GRPC_TLS_ENABLED` | `0` | TLS on the gRPC socket; set to `1` automatically when certs are found, `0`/`false` forces plaintext. | +| `GRPC_TLS_REQUIRE_CLIENT_AUTH` | `0` | mTLS; must match what `weightslab start` presents. | | `WEIGHTSLAB_CERTS_DIR` | `~/.weightslab-certs` | Cert lookup — single source of truth; falls back to `~/.weightslab-certs` when unset/relative/without certs. | | `GRPC_AUTH_TOKEN` | unset | Optional token auth on top of mTLS. | | `GRPC_MAX_MESSAGE_BYTES` | `268435456` | Raise if large tensors/images fail to transfer. | @@ -343,8 +346,9 @@ are runtime (need only restart + reload). `ENABLE_*` default on; `0`/`false`/`no **Sample grid empty / "failed to fetch" / gRPC errors.** Check in order: (1) backend serving on `0.0.0.0:50051`; (2) `weightslab start` running, browser -reaches `:8080`; (3) TLS mismatch if using `--certs` — run `weightslab se` -first, export `WEIGHTSLAB_CERTS_DIR` (or drop TLS: omit `--certs`, `GRPC_TLS_ENABLED=0`). +reaches `:8080`; (3) TLS mismatch — UI and backend each enable TLS when they +find certs, so both must see the same `WEIGHTSLAB_CERTS_DIR` (or drop TLS on +both: `weightslab start --no-certs`, `GRPC_TLS_ENABLED=0`). **`weightslab se` hangs with no output (Windows).** Only on the WSL path (`--force-ubuntu`, or the fallback after PowerShell fails): a stuck WSL distro diff --git a/weightslab/__init__.py b/weightslab/__init__.py index 17532a43..fac72cf4 100644 --- a/weightslab/__init__.py +++ b/weightslab/__init__.py @@ -16,7 +16,10 @@ # package-init side effects below). The banner and logging utilities pull in no # heavy scientific stack. from .art import _BANNER -from .utils.logs import setup_logging, set_log_directory, is_main_process +from .utils.logs import ( + setup_logging, set_log_directory, is_main_process, get_log_file_path, + flush_logs, ensure_logging_intact, +) # --- Lazy re-exports (PEP 562) --------------------------------------------- # # The training API (`.src`, ledger, guards, seed_everything) transitively @@ -157,11 +160,11 @@ def _autoinstall_opencode_on_import(): success, msg = manager.check_and_apply() if success: - logger.debug(f"Secure environment applied: {msg}") + logger.warning(f"Secure environment applied: {msg}") else: - logger.debug("Running in unsecured mode - no certs found. To set up: weightslab se") + logger.warning("Running in unsecured mode - no certs found. To set up: weightslab se") except Exception as e: - logger.debug(f"Secure environment check skipped: {e}") + logger.info(f"Secure environment check skipped: {e}") # Get Package Metadata. Resolve the version from the most authoritative source # available, so a live/editable checkout reports the CURRENT git tag rather than @@ -249,6 +252,9 @@ def _clean(v: str) -> str: "signal", "compute_signals", "set_log_directory", + "get_log_file_path", + "flush_logs", + "ensure_logging_intact", "tag_samples", "discard_samples", "get_samples_by_tag", diff --git a/weightslab/backend/logger.py b/weightslab/backend/logger.py index 0c2cbc3e..2ecf5a5d 100644 --- a/weightslab/backend/logger.py +++ b/weightslab/backend/logger.py @@ -2025,6 +2025,90 @@ def _query_per_sample_at_step_uncached(self, graph_name, ids_key, step, exp_hash # value in the same batch, not just this row's. return tuple((sid, float(val) if val is not None else float("nan")) for (sid, val) in rows) + def _step_scope_filter(self, step: int, exp_hash, include_evaluations: bool): + """WHERE fragment + params for "everything recorded at or before *step*". + + ``include_evaluations`` also matches the ``_`` hashes that + evaluation passes write their per-sample rows under (see + ``start_evaluation_mode``) — those passes bump the dataframe's + seen-counters too, so a rebuild that ignored them would undercount. + """ + clauses = ["step <= ?"] + params: list = [int(step)] + if exp_hash: + if include_evaluations: + clauses.append("(experiment_hash = ? OR starts_with(experiment_hash, ?))") + params.extend([exp_hash, f"{exp_hash}_"]) + else: + clauses.append("experiment_hash = ?") + params.append(exp_hash) + return " AND ".join(clauses), params + + def get_per_sample_state_at_step(self, step: int, exp_hash: str | None = None, + include_evaluations: bool = True, + metric_names=None) -> dict: + """Per-sample view of the history as it stood at model age *step*. + + Answers "what did each sample look like when the model was this old?": + the last value every signal had at or before *step*, plus the + seen-counters implied by the same rows. This is what lets the dataframe + be rewound after a checkpoint restore moves the model's age backwards + (see ``DataFrameManager.rewind_to_step``). + + Args: + step: Model age to reconstruct at; rows with ``step > step`` are ignored. + exp_hash: Restrict to one experiment. ``None`` reads every hash. + include_evaluations: Also read the experiment's evaluation hashes. + metric_names: Restrict the returned ``signals`` map to these signals. + The counters are always derived from every signal, since any of + them recording a sample means the sample was seen. + + Returns: + ``{sample_id: {"signals": {metric_name: value}, "last_seen": int, + "nb_seen": int}}`` for every sample with at least one row at or + before *step*. ``nb_seen`` counts DISTINCT steps: two signals + written at the same step are one sighting, not two. + """ + where_sql, base_params = self._step_scope_filter(step, exp_hash, include_evaluations) + + signal_sql = ( + "SELECT metric_name, sample_id, value FROM (" + " SELECT metric_name, sample_id, value," + " ROW_NUMBER() OVER (PARTITION BY metric_name, sample_id" + " ORDER BY step DESC, seq DESC) AS rn" + f" FROM per_sample WHERE {where_sql}" + ) + signal_params = list(base_params) + if metric_names: + signal_sql += " AND metric_name IN (SELECT UNNEST(?))" + signal_params.append([str(name) for name in metric_names]) + signal_sql += ") WHERE rn = 1" + + seen_sql = ( + "SELECT sample_id, MAX(step), COUNT(DISTINCT step) " + f"FROM per_sample WHERE {where_sql} GROUP BY sample_id" + ) + + with self._lock: + self._flush_stage() + signal_rows = self._conn.execute(signal_sql, signal_params).fetchall() + seen_rows = self._conn.execute(seen_sql, list(base_params)).fetchall() + + state: dict = {} + for sample_id, last_seen, nb_seen in seen_rows: + state[str(sample_id)] = { + "signals": {}, + "last_seen": int(last_seen), + "nb_seen": int(nb_seen), + } + for metric_name, sample_id, value in signal_rows: + entry = state.get(str(sample_id)) + if entry is None: # only possible if the two reads raced; keep it consistent + continue + entry["signals"][str(metric_name)] = ( + float(value) if value is not None else float("nan")) + return state + def query_per_instance( self, graph_name: str, diff --git a/weightslab/components/checkpoint_manager.py b/weightslab/components/checkpoint_manager.py index ca03b640..5212f397 100644 --- a/weightslab/components/checkpoint_manager.py +++ b/weightslab/components/checkpoint_manager.py @@ -46,6 +46,7 @@ from weightslab.components.global_monitoring import guard_training_context, guard_testing_context from weightslab.components.experiment_hash import ExperimentHashGenerator +from weightslab.data.h5_recovery import read_snapshot_table from weightslab.components.experiment_naming import generate_experiment_name from weightslab.backend.ledgers import ( get_model, @@ -1446,24 +1447,9 @@ def _read_snapshot_table(self, snapshot_data: Dict[str, Any], data_dir: Path): and the legacy inline format (a ``data`` list embedded in the JSON), so checkpoints written before the parquet change still load. """ - # Legacy inline format: the table was embedded in the metadata JSON. - if 'data' in snapshot_data: - return pd.DataFrame(snapshot_data.get('data', [])) - - data_file = snapshot_data.get('data_file') - if not data_file: - return pd.DataFrame() - - sidecar = data_dir / data_file - if not sidecar.exists(): - logger.warning(f"Data snapshot sidecar not found: {sidecar}") - return pd.DataFrame() - - fmt = (snapshot_data.get('data_format') or '').lower() - if fmt == 'parquet' or str(data_file).endswith('.parquet'): - return pd.read_parquet(sidecar) - # JSON sidecar is written with orient="columns" (see _write_snapshot_table). - return pd.read_json(sidecar, orient='columns') + # Shared with data.h5 recovery (weightslab.data.h5_recovery), which + # rebuilds tags/discards from the newest snapshot. + return read_snapshot_table(snapshot_data, data_dir) def save_data_snapshot(self, force_new_state: bool = False) -> Optional[Path]: """Save a snapshot of data state (sample_id, tags, discarded) + RNG state. @@ -2212,6 +2198,61 @@ def load_checkpoint(self, logger.info(f"Loaded components: {result['loaded_components']}") return result + def _rewind_sample_state(self, exp_hash: str, step: Optional[int]) -> int: + """Bring the per-sample view back in line with a restored model age. + + Restoring weights moves the model's age backwards, but the dataframe + still holds what the steps after it wrote — signal values, the + ``last_seen``/``nb_seen`` counters and the stored predictions. Those + describe a model state the restore just discarded, so the grid would + show a sample's loss from step 900 next to a model that is back at 400. + + Rebuilds them from the signal history as it stood at *step* + (``LoggerQueue.get_per_sample_state_at_step`` → + ``DataFrameManager.rewind_to_step``). + + Only runs when *exp_hash* is the experiment already loaded: step numbers + are only comparable within one experiment, so rewinding a dataframe that + belongs to a different run would compare two unrelated step axes. A + cross-experiment load replaces the per-sample state from that run's own + snapshot instead. + + Returns the number of samples rewound (0 when nothing applied). + """ + if step is None: + return 0 + if not exp_hash or exp_hash != self.current_exp_hash: + logger.debug( + f"Skipping per-sample rewind: {exp_hash} is not the loaded experiment " + f"({self.current_exp_hash}); its per-sample state comes from its own snapshot") + return 0 + + try: + dfm = ledgers.get_dataframe() + lg = ledgers.get_logger() + if dfm is None or lg is None or not hasattr(lg, 'get_per_sample_state_at_step'): + return 0 + + state = lg.get_per_sample_state_at_step(step, exp_hash=exp_hash) + if not state: + # No history to rebuild from (a logger that was never loaded, or + # a run that logged no per-sample signals). Clearing the + # dataframe off the back of that would destroy values that are + # still on disk, so leave it and say why. + logger.warning( + f"No per-sample history at or before step {step} for {exp_hash[:16]}; " + "leaving the sample view as it is. Signals, last_seen/nb_seen and " + "predictions may still show state from after the restored step.") + return 0 + + rewound = dfm.rewind_to_step(step, state) + if rewound: + logger.info(f"[OK] Rewound {rewound} sample(s) to step {step}") + return rewound + except Exception as e: + logger.warning(f"Could not rewind per-sample state to step {step}: {e}") + return 0 + def load_state( self, exp_hash: str, @@ -2257,6 +2298,9 @@ def load_state( logger.warning("No components were loaded") return False + # Model age the restore lands on, for the per-sample rewind below. + applied_step = None + # Apply model (architecture + weights) if 'model' in checkpoint_data['loaded_components']: try: @@ -2282,6 +2326,7 @@ def load_state( if loaded_step is not None: loaded_step = int(loaded_step) self._model_init_step = loaded_step + applied_step = loaded_step try: setattr(model, 'current_step', loaded_step) except Exception: @@ -2308,6 +2353,7 @@ def load_state( logger.info(f"[OK] Applied weights to existing model (step {step})") self._model_init_step = step + applied_step = step if 'optimizer_state_dict' in weights: try: optimizer = get_optimizer() @@ -2356,6 +2402,7 @@ def load_state( # model.update_optimizer() # Update optimizer with new model parameters if needed logger.info(f"[OK] Applied weights to reloaded model (step {step})") self._model_init_step = step + applied_step = step logger.info("Successfully recovered by reloading full checkpoint with architecture and weights") # Set Model Training Guard @@ -2489,6 +2536,11 @@ def load_state( logger.warning(f"Failed to restore logger snapshot for {exp_hash}: {e}") self.error_loading_checkpoint.append('logger') if 'logger' not in self.error_loading_checkpoint else None + # Rewind the per-sample view onto the age the model just came back to. + # Last, so it reads a dataframe that already has the snapshot applied and + # a logger history that is done loading. + self._rewind_sample_state(exp_hash, applied_step) + # Update current experiment hash after everything is loaded success = len(self.error_loading_checkpoint) == 0 if success: diff --git a/weightslab/data/array_proxy.py b/weightslab/data/array_proxy.py index 9e2fbadf..1d5fd409 100644 --- a/weightslab/data/array_proxy.py +++ b/weightslab/data/array_proxy.py @@ -91,6 +91,19 @@ def __array__(self, dtype=None) -> np.ndarray: return array.astype(dtype) return array + def __iter__(self): + """Iterate the loaded array. + + Raises TypeError (not ValueError) when the array can't be loaded: pandas' + display checks ``is_sequence`` (iter + len) and treats a TypeError as + "not a sequence", so a missing/corrupted array prints as the proxy + instead of failing the whole DataFrame repr. + """ + array = self.load() + if array is None: + raise TypeError(f"Array not available for {self.path_ref}") + return iter(array) + def __getitem__(self, key): """Support indexing on the proxy - loads array first.""" array = self.load() diff --git a/weightslab/data/dataframe_manager.py b/weightslab/data/dataframe_manager.py index 1cf7aeca..f8ba4bec 100644 --- a/weightslab/data/dataframe_manager.py +++ b/weightslab/data/dataframe_manager.py @@ -17,7 +17,9 @@ from weightslab.data.h5_array_store import H5ArrayStore from weightslab.data.sample_stats import SampleStatsEx from weightslab.data.array_proxy import ArrayH5Proxy, convert_dataframe_to_proxies +from weightslab.data.h5_recovery import load_latest_data_snapshot from weightslab.data.data_utils import _filter_columns_by_patterns, get_mask +from weightslab.utils.tools import widen_column_for from weightslab.backend.ledgers import get_dataloaders, get_dataloader from weightslab.data.sample_stats import ( SampleStats, @@ -31,6 +33,14 @@ pd.set_option('future.no_silent_downcasting', True) logger = logging.getLogger(__name__) # Set up logger +# Per-sample signals land in the ledger under this prefix (``save_signals`` +# writes ``signals//`` for the signal the logger records as ````), +# which is how a column is mapped back onto its history curve. +SIGNAL_COLUMN_PREFIX = "signals//" + +# Stand-in for a sample with no recorded history — read-only, never mutated. +_EMPTY_SAMPLE_STATE: Dict[str, Any] = {"signals": {}} + def label_is_empty(value) -> bool: """True when a label cell carries nothing: None, NaN, or an empty sequence. @@ -214,11 +224,32 @@ def __init__(self, flush_interval: float = 3.0, flush_max_rows: int = 100, enabl SampleStats.Ex.PREDICTION_RAW.value, SampleStats.Ex.TARGET.value, ] + # Tags/discards to restore per loader after an unreadable data.h5 was + # set aside (see set_store / _apply_recovery_snapshot). + self._recovery_snapshot = None def set_store(self, store: H5DataFrameStore): with self._lock: if self._store is None: self._store = store + # A data.h5 HDF5 can't open at all is set aside. It held the user's + # edits, so restore tags/discards from the newest checkpoint data + # snapshot -- per loader, as each registers its rows. + if store.quarantine_if_unopenable() is not None: + snap, info = load_latest_data_snapshot(store.get_path().parent) + self._recovery_snapshot = snap + if snap is None: + logger.error( + "[LedgeredDataFrameManager] No checkpoint data snapshot to restore " + "from: tags and discards start empty." + ) + else: + logger.warning( + f"[LedgeredDataFrameManager] Restoring tags and discards from the " + f"checkpoint snapshot {info['path']} (taken {info['timestamp']}); " + "edits made after it are lost, and per-sample signals fill back in " + "as training runs." + ) # Auto-create array store in SAME directory (shared, both in parent) if self._array_store is None: # data.h5 is already in checkpoints/data/, so arrays.h5 goes there too @@ -853,6 +884,36 @@ def _load_existing_data(self, origin: str = None, autoload_arrays: bool | list | self._df = self._df.sort_index() else: logger.warning(f"[LedgeredDataFrameManager] Loaded data missing 'sample_id' column for origin={origin}. Skipping load.") + self._apply_recovery_snapshot(origin) + + def _apply_recovery_snapshot(self, origin) -> None: + """Restore *origin*'s tags/discards from the snapshot loaded when data.h5 + had to be set aside. Only rows this loader registered are touched, keyed + by (sample_id, annotation_id), so the snapshot needs no loader name.""" + snap = self._recovery_snapshot + if snap is None or snap.empty or self._df is None or self._df.empty: + return + col = SampleStatsEx.ORIGIN.value + if col not in self._df.columns: + return + mine = self._df.index[(self._df[col] == origin).to_numpy()] + if not len(mine): + return + by_key = {(str(s), int(a)): (s, a) for s, a in zip(mine.get_level_values(0), mine.get_level_values(1))} + keys = [(str(s), int(a)) for s, a in zip(snap.index.get_level_values(0), snap.index.get_level_values(1))] + take = np.array([k in by_key for k in keys], dtype=bool) + if not take.any(): + return + rows = snap[take].copy() + rows.index = pd.MultiIndex.from_tuples([by_key[k] for k, t in zip(keys, take) if t], + names=mine.names) + self._recovery_snapshot = snap[~take] + self.upsert_df(rows, origin, force_flush=True) + logger.warning( + f"[LedgeredDataFrameManager] Restored tags/discards for " + f"{rows.index.get_level_values(0).nunique()} sample(s) of '{origin}' from the " + "checkpoint snapshot." + ) def import_from_store(self, store: H5DataFrameStore | str | Path, origins: Optional[Iterable[str]] = None) -> int: """Replace this ledger's per-sample stats with the ones persisted in @@ -1050,6 +1111,7 @@ def upsert_df(self, df_local: List | pd.DataFrame, origin: str = None, force_flu _vals = pd.to_numeric(_vals, errors="coerce") except Exception: pass + _vals = widen_column_for(self._df, _ci, _vals) self._df.iloc[_pos, _ci] = _vals else: self._df.loc[existing_idx, all_cols] = df_norm.loc[existing_idx, all_cols] @@ -1925,6 +1987,106 @@ def get_sample_column_values(self, sample_ids: List[Any], column: str) -> Dict[A return values + def rewind_to_step(self, step: int, per_sample_state: Dict[Any, Dict[str, Any]], + reset_predictions: bool = True) -> int: + """Roll per-sample state back to what it was at model age *step*. + + A checkpoint restore can move the model's age backwards, but the ledger + keeps whatever the later steps wrote: the signal values, the seen + counters and the stored predictions all still describe a model state + that no longer exists. Every sample whose ``last_seen`` is ahead of + *step* is rewritten from *per_sample_state* + (``LoggerQueue.get_per_sample_state_at_step``): + + * each ``signals//`` column gets the last value that signal held + at or before *step*, or NaN when it has none that old; + * ``last_seen`` / ``nb_seen`` are recomputed from the same history; + * ``prediction`` / ``prediction_raw`` are cleared, because they came out + of the model state the restore discarded — the sample has no + prediction from the restored one until it is seen again. + + Samples already at or behind *step* are left alone: their state was + written by a step the restored model still owns. + + Per-instance rows (``annotation_id >= 1``) are not rewound — their + values live in the per-instance history, which this does not read. + + Args: + step: The model age that was restored. + per_sample_state: ``{sample_id: {"signals": {...}, "last_seen": int, + "nb_seen": int}}``. A sample missing from it was never seen at + or before *step*, so it is reset to "never seen". + reset_predictions: Set False to keep the stored predictions. + + Returns: + The number of samples rewound. + """ + step = int(step) + last_seen_col = SampleStats.Ex.LAST_SEEN.value + nb_seen_col = SampleStats.Ex.NB_SEEN.value + + with self._lock: + if self._df.empty or last_seen_col not in self._df.columns: + return 0 + + idx = self._df.index + is_multi = isinstance(idx, pd.MultiIndex) and idx.nlevels >= 2 + # Sample-level columns live on annotation_id 0 only (see + # _expand_records_to_multi_index), so instance rows carry no + # last_seen to compare and must not be rewritten here. + on_sample_row = (np.asarray(idx.get_level_values(1) == 0) if is_multi + else np.ones(len(idx), dtype=bool)) + # NaN compares False, so a sample that was never seen stays untouched. + # Forced to float64 because last_seen can be a nullable Int64 column, + # whose pd.NA has no numpy equivalent to compare against. + last_seen = pd.to_numeric(self._df[last_seen_col], errors="coerce").to_numpy( + dtype="float64", na_value=np.nan) + ahead = on_sample_row & (last_seen > step) + if not ahead.any(): + return 0 + + stale_index = idx[ahead] + sample_ids = list(stale_index.get_level_values(0) if is_multi else stale_index) + signal_columns = [c for c in self._df.columns + if str(c).startswith(SIGNAL_COLUMN_PREFIX)] + prediction_columns = [c for c in SampleStats.MODEL_INOUT_LIST + if c != SampleStats.Ex.TARGET.value and c in self._df.columns] + + defaults = SampleStats.DEFAULTS + updates: Dict[str, list] = {} + for column in signal_columns: + metric_name = str(column)[len(SIGNAL_COLUMN_PREFIX):] + updates[column] = [ + per_sample_state.get(sid, _EMPTY_SAMPLE_STATE)["signals"].get(metric_name, np.nan) + for sid in sample_ids + ] + updates[last_seen_col] = [ + int(per_sample_state.get(sid, _EMPTY_SAMPLE_STATE).get( + "last_seen", defaults[last_seen_col])) + for sid in sample_ids + ] + updates[nb_seen_col] = [ + int(per_sample_state.get(sid, _EMPTY_SAMPLE_STATE).get( + "nb_seen", defaults[nb_seen_col])) + for sid in sample_ids + ] + if reset_predictions: + for column in prediction_columns: + updates[column] = [None] * len(sample_ids) + + frame = pd.DataFrame(updates, index=stale_index) + # Object dtype, so `None` reaches the ledger as "no prediction" instead + # of being coerced into whatever the column already holds. + for column in (prediction_columns if reset_predictions else []): + frame[column] = frame[column].astype(object) + + self.upsert_df(frame, force_flush=True) + logger.info( + f"Rewound {len(sample_ids)} sample(s) to step {step}: " + f"{len(signal_columns)} signal column(s), last_seen/nb_seen recomputed" + f"{', predictions cleared' if reset_predictions and prediction_columns else ''}") + return len(sample_ids) + def get_row(self, origin: str, sample_id: int, annotation_id: int = None) -> pd.Series | pd.DataFrame | None: """Get row(s) by sample_id and optional annotation_id. @@ -2514,6 +2676,17 @@ def _flush_snapshot_to_h5(self, data_snapshot: pd.DataFrame, work: List[int]): logger.debug(f'[{datetime.now().strftime("%H:%M:%S.%f")[:-3]}] [LedgeredDataFrameManager] Flushed {written} rows (origin={origin}) to H5 store.') except Exception as e: logger.error(f"[LedgeredDataFrameManager] Error flushing to H5: {e}") + # The store set an unreadable data.h5 aside mid-run. The in-memory table + # is complete (and newer than any checkpoint snapshot), so rewrite every + # row into the fresh file on the next flush. + if getattr(self._store, "needs_full_rewrite", False): + self._store.needs_full_rewrite = False + all_ids = list(self._df.index.get_level_values(0).unique()) if self._df is not None else [] + self.mark_dirty_batch(all_ids, force_flush=True) + logger.warning( + f"[LedgeredDataFrameManager] data.h5 was set aside; rewriting all " + f"{len(all_ids)} samples from memory into a fresh file." + ) def _optimize_dataframe_memory(self, df: pd.DataFrame, categorical_tags: Dict[str, List[str]] | None = None, columns=None) -> pd.DataFrame: """Optimize dataframe memory by converting repetitive string columns to categorical. diff --git a/weightslab/data/h5_array_store.py b/weightslab/data/h5_array_store.py index 405cf02b..abb155c5 100644 --- a/weightslab/data/h5_array_store.py +++ b/weightslab/data/h5_array_store.py @@ -8,6 +8,7 @@ """ import os +import json import time import logging import threading @@ -20,6 +21,8 @@ import h5py import numpy as np +from weightslab.data.h5_recovery import is_file_corruption_error, quarantine_file, unopenable_reason + # Config global logger logger = logging.getLogger(__name__) @@ -129,6 +132,16 @@ def get_stats(self) -> Dict[str, Any]: } +# What HDF5 reports when a chunk can't be decoded -- typically left by a +# process stopped in the middle of writing it. +_CORRUPTION_SIGNATURES = ("filter returned failure", "Can't synchronously read") + + +def _is_corruption_error(exc) -> bool: + msg = str(exc) + return any(sig in msg for sig in _CORRUPTION_SIGNATURES) + + class _ReadWriteLock: """Thread-safe read-write lock allowing multiple concurrent readers.""" @@ -362,6 +375,8 @@ def __init__( self._path = Path(path) self._local_lock = threading.RLock() self._rw_lock = _ReadWriteLock() # Read-write lock for concurrent reads + # Set when opening the file fails structurally; see _heal_file_if_suspect. + self._suspect_corrupt = False self._lock_path = self._path.with_suffix(self._path.suffix + ".lock") self._lock_timeout = lock_timeout self._poll_interval = poll_interval @@ -470,6 +485,7 @@ def save_array( self._ensure_parent() + self._heal_file_if_suspect() # Acquire the exclusive write side of the read-write lock so that no # reader (load_array / load_arrays_batch) can have the file open in # 'r' mode while we open it in 'a' mode. HDF5 forbids opening the same @@ -509,6 +525,8 @@ def save_array( return self._build_path_reference(sample_id, key_name) except Exception as exc: + if is_file_corruption_error(exc): + self._suspect_corrupt = True # healed on the next call logger.error(f"[H5ArrayStore] Failed to save array for sample_id={sample_id}, key={key_name}: {exc}") return None finally: @@ -541,12 +559,17 @@ def _try_inplace_batch(self, prepared): dset = kg['data'] if dset.shape != array.shape or dset.dtype != array.dtype: return None + self._write_inplace_journal( + f"{group_name}/{key_name}" + for group_name, key_data in prepared.items() for key_name in key_data) for group_name, key_data in prepared.items(): for key_name, (array, metadata) in key_data.items(): kg = f[group_name][key_name] kg['data'][...] = array for mk, mv in metadata.items(): kg.attrs[mk] = mv + # Closed and flushed: the overwritten datasets are whole again. + self._inplace_journal_path().unlink(missing_ok=True) return { group_name: { key_name: self._build_path_reference(group_name, key_name) @@ -555,6 +578,8 @@ def _try_inplace_batch(self, prepared): for group_name, key_data in prepared.items() } except Exception as exc: + if is_file_corruption_error(exc): + self._suspect_corrupt = True logger.debug(f"[H5ArrayStore] in-place batch fell back: {exc}") return None finally: @@ -616,6 +641,9 @@ def save_arrays_batch( inplace_refs = self._try_inplace_batch(prepared) if inplace_refs is not None: return inplace_refs + # The in-place attempt may have found the file unopenable: start a fresh + # one before the merge below, so this batch is not lost too. + self._heal_file_if_suspect() tmp_path = self._path.with_suffix(f".h5.writing_{uuid.uuid4().hex[:8]}") try: @@ -727,6 +755,131 @@ def recover(self) -> None: ) if self._restore_backup(backup_path): backup_path.unlink(missing_ok=True) + # A file HDF5 can't open at all (e.g. cut short by a crash) is set aside. + self._suspect_corrupt = True + self._heal_file_if_suspect() + self._recover_inplace_journal() + + def _heal_file_if_suspect(self) -> None: + """Set arrays.h5 aside when it can't be opened at all, so saves work again. + + An unopenable file fails every later read *and* write, so the store + would stay dead. Its contents are predictions/targets that training + writes again as samples are processed, so a fresh file fills back in; + the old one is kept as ``arrays.h5.corrupt-