tsseg / src /params.py
fchavelli's picture
fix(app): neutral coloring when no CP found + editable n_segments for CLaSP/Clap (0 = auto)
7586faf
Raw History Blame Contribute Delete
19.5 kB
"""Dynamic parameter rendering for Gradio based on tsseg param_schema + YAML fallback.
Priority order for each detector:
1. ``tsseg.algorithms.param_schema.get_ui_hints`` (declarative schema with
data-dependent bounds).
2. YAML tunable ranges in ``src/config/<detector>.yaml``.
3. Pure introspection of the ``__init__`` signature (legacy fallback).
When a *data context* is provided (``n_samples``, ``n_channels``), data-
dependent constraints are resolved and default values that violate them are
**automatically clamped** with a warning.
"""
from __future__ import annotations
import inspect
import logging
from pathlib import Path
from typing import Any, Dict, List, Optional, Tuple
import yaml
logger = logging.getLogger(__name__)
_CONFIG_DIR = Path(__file__).parent / "config"
_YAML_CACHE: Dict[str, dict] = {}
# ---------------------------------------------------------------------------
# Data context helper
# ---------------------------------------------------------------------------
def build_data_context(
signal: Any,
*,
n_channels: int | None = None,
) -> Dict[str, Any]:
"""Build a ``data_ctx`` dict from a signal array.
Parameters
----------
signal : array-like
The time-series array, shape ``(n_samples,)`` or ``(n_samples, n_dims)``.
n_channels : int, optional
Override channel count (inferred from *signal* otherwise).
Returns
-------
dict
``{"n_samples": …, "n_channels": …}``
"""
import numpy as np
sig = np.asarray(signal)
n_samples = sig.shape[0]
if n_channels is None:
n_channels = sig.shape[1] if sig.ndim > 1 else 1
return {"n_samples": n_samples, "n_channels": n_channels}
# ---------------------------------------------------------------------------
# YAML config (unchanged, used as fallback / override)
# ---------------------------------------------------------------------------
def _load_yaml_config(detector_name: str) -> Optional[dict]:
"""Load YAML config with tunable parameter ranges for a detector."""
if detector_name in _YAML_CACHE:
return _YAML_CACHE[detector_name]
candidates = [
detector_name.lower().replace("detector", ""),
detector_name.lower().replace("detector", "").replace("_", "-"),
detector_name.lower(),
]
for candidate in candidates:
path = _CONFIG_DIR / f"{candidate}.yaml"
if path.exists():
try:
with open(path) as f:
config = yaml.safe_load(f)
_YAML_CACHE[detector_name] = config
return config
except Exception as e:
logger.warning(f"Failed to load YAML config {path}: {e}")
_YAML_CACHE[detector_name] = None
return None
def get_tunable_ranges(detector_name: str) -> Dict[str, list]:
"""Get tunable parameter ranges from YAML config."""
config = _load_yaml_config(detector_name)
if config is None:
return {}
tunable = config.get("tunable_parameters", [])
ranges: Dict[str, list] = {}
if isinstance(tunable, list):
for item in tunable:
if isinstance(item, dict):
ranges.update(item)
return ranges
# ---------------------------------------------------------------------------
# Schema-aware parameter info
# ---------------------------------------------------------------------------
def _has_parameter_schema(cls: type) -> bool:
"""Return True if *cls* declares a non-empty ``_parameter_schema``."""
schema = getattr(cls, "_parameter_schema", None)
return bool(schema)
def _clamp_default(
value: Any,
lo: float | int | None,
hi: float | int | None,
ptype: str,
) -> Tuple[Any, bool]:
"""Clamp *value* to ``[lo, hi]``. Returns ``(clamped, was_adjusted)``."""
if value is None:
return value, False
try:
v = float(value)
except (TypeError, ValueError):
return value, False
adjusted = False
if lo is not None and v < lo:
v = lo
adjusted = True
if hi is not None and v > hi:
v = hi
adjusted = True
if adjusted:
v = int(v) if ptype == "int" else v
return v, adjusted
def get_param_info(
cls: type,
data_ctx: Dict[str, Any] | None = None,
) -> List[Dict[str, Any]]:
"""Extract parameter info for UI rendering.
When a ``_parameter_schema`` is declared on *cls*, it is preferred over
the legacy YAML / introspection path. In both cases, YAML tunable
ranges are merged in to refine the choices.
Parameters
----------
cls : type
Detector class.
data_ctx : dict, optional
Data context (``{"n_samples": …, "n_channels": …}``).
Returns
-------
list[dict]
Per-parameter dicts with keys: ``name``, ``type``, ``default``,
``min``, ``max``, ``step``, ``choices``, ``description``, ``group``,
``nullable``, ``hidden``, ``adjusted``, ``adjust_reason``.
"""
if _has_parameter_schema(cls):
result = _param_info_from_schema(cls, data_ctx)
else:
result = _param_info_legacy(cls)
return apply_param_overrides(cls.__name__, result)
# ---------------------------------------------------------------------------
# Schema-based path (new)
# ---------------------------------------------------------------------------
def _param_info_from_schema(
cls: type,
data_ctx: Dict[str, Any] | None = None,
) -> List[Dict[str, Any]]:
"""Build param info list from ``get_ui_hints``."""
from tsseg.algorithms.param_schema import get_ui_hints
hints = get_ui_hints(cls, data_ctx=data_ctx)
if not hints:
return _param_info_legacy(cls)
detector_name = cls.__name__
tunable_ranges = get_tunable_ranges(detector_name)
params: List[Dict[str, Any]] = []
for pname, h in hints.items():
if h.get("hidden", False):
continue
ptype = h.get("type", "any")
default = h.get("default")
lo = h.get("min")
hi = h.get("max")
nullable = h.get("nullable", False)
choices = h.get("choices")
description = h.get("description", "")
group = h.get("group", "")
info: Dict[str, Any] = {
"name": pname,
"default": default,
"description": description,
"group": group,
"nullable": nullable,
"adjusted": False,
"adjust_reason": "",
}
# --- Merge YAML tunable ranges if present --------------------------
if pname in tunable_ranges and choices is None:
yaml_choices = tunable_ranges[pname]
if isinstance(yaml_choices, list) and all(isinstance(c, str) for c in yaml_choices):
choices = yaml_choices
# --- Determine UI widget type --------------------------------------
if choices:
info["type"] = "choice"
info["choices"] = choices
if default is not None and default not in choices:
choices.insert(0, default)
elif ptype == "bool":
info["type"] = "bool"
elif ptype == "int":
info["type"] = "int"
effective_lo = lo if lo is not None else 0
effective_hi = hi if hi is not None else _fallback_upper(default, effective_lo)
info["min"] = int(effective_lo)
info["max"] = int(effective_hi)
info["step"] = 1
# Auto-clamp default
if default is not None:
clamped, was = _clamp_default(default, effective_lo, effective_hi, "int")
if was:
info["default"] = clamped
info["adjusted"] = True
info["adjust_reason"] = (
f"Default {pname}={default} adjusted to {clamped} (series of {data_ctx['n_samples']} points)"
if data_ctx
else f"Default {pname}={default} out of bounds, adjusted to {clamped}"
)
elif ptype == "float":
info["type"] = "float"
effective_lo = lo if lo is not None else 0.0
effective_hi = hi if hi is not None else _fallback_upper_f(default, effective_lo)
info["min"] = float(effective_lo)
info["max"] = float(effective_hi)
info["step"] = round((effective_hi - effective_lo) / 100, 6) or 0.01
# Auto-clamp default
if default is not None:
clamped, was = _clamp_default(default, effective_lo, effective_hi, "float")
if was:
info["default"] = clamped
info["adjusted"] = True
info["adjust_reason"] = (
f"Default {pname}={default} adjusted to {clamped} (series of {data_ctx['n_samples']} points)"
if data_ctx
else f"Default {pname}={default} out of bounds, adjusted to {clamped}"
)
elif ptype == "str":
info["type"] = "str"
else:
info["type"] = "text"
params.append(info)
return params
def _fallback_upper(default: Any, lo: float | int) -> int:
"""Heuristic upper bound when the schema has no finite upper bound."""
if default is not None:
try:
d = int(default)
return max(d * 3 + 10, int(lo) + 10)
except (TypeError, ValueError):
pass
return int(lo) + 100
def _fallback_upper_f(default: Any, lo: float) -> float:
"""Heuristic upper bound for floats."""
if default is not None:
try:
d = float(default)
span = max(abs(d), 0.1)
return round(d + 3 * span, 6)
except (TypeError, ValueError):
pass
return lo + 10.0
# ---------------------------------------------------------------------------
# Per-detector parameter overrides
# ---------------------------------------------------------------------------
# Some upstream detectors accept a *sentinel* string (e.g. ``"learn"``) on
# parameters that would otherwise be integers. The introspection-based UI
# would render them as plain text boxes, which is both unfriendly and
# error-prone (typing ``"5"`` keeps the value as a string and breaks the call).
#
# This map injects friendlier widgets for those parameters; the demo then
# translates the sentinel value (``0``) back to the upstream string in
# :func:`apply_sentinel_translations` before instantiating the detector.
_PARAM_OVERRIDES: Dict[str, Dict[str, Dict[str, Any]]] = {
"ClaspDetector": {
"n_segments": {
"type": "int",
"default": 0,
"min": 0,
"max": 50,
"step": 1,
"description": "Number of segments to detect. 0 = auto (learn from data).",
},
},
"ClapDetector": {
"n_segments": {
"type": "int",
"default": 0,
"min": 0,
"max": 50,
"step": 1,
"description": "Number of segments to detect. 0 = auto (learn from data).",
},
"n_change_points": {
"type": "int",
"default": 0,
"min": 0,
"max": 50,
"step": 1,
"description": "Number of change points (overrides n_segments). 0 = auto.",
},
},
}
# Sentinel value used in the UI to mean "fall back to the upstream auto mode".
_AUTO_SENTINEL = 0
_AUTO_REPLACEMENT: Dict[str, Dict[str, Any]] = {
"ClaspDetector": {"n_segments": "learn"},
"ClapDetector": {"n_segments": None, "n_change_points": None},
}
def apply_param_overrides(detector_name: str, params: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
"""Replace the UI hint of selected params with a friendlier widget."""
overrides = _PARAM_OVERRIDES.get(detector_name)
if not overrides:
return params
out: List[Dict[str, Any]] = []
for p in params:
if p["name"] in overrides:
merged = dict(p)
merged.update(overrides[p["name"]])
# Reset adjustment flags since we replaced the spec entirely
merged.setdefault("group", "")
merged.setdefault("nullable", False)
merged.setdefault("adjusted", False)
merged.setdefault("adjust_reason", "")
out.append(merged)
else:
out.append(p)
return out
def apply_sentinel_translations(detector_name: str, params: Dict[str, Any]) -> Dict[str, Any]:
"""Translate UI sentinel values (``0`` = auto) back to upstream defaults."""
rules = _AUTO_REPLACEMENT.get(detector_name)
if not rules:
return params
out = dict(params)
for pname, replacement in rules.items():
if pname in out and out[pname] == _AUTO_SENTINEL:
out[pname] = replacement
return out
# ---------------------------------------------------------------------------
# Legacy path (introspection + YAML, no param_schema)
# ---------------------------------------------------------------------------
def _param_info_legacy(cls: type) -> List[Dict[str, Any]]:
"""Original introspection-based param extraction (unchanged logic)."""
sig = inspect.signature(cls.__init__)
detector_name = cls.__name__
tunable_ranges = get_tunable_ranges(detector_name)
params: List[Dict[str, Any]] = []
for param in sig.parameters.values():
if param.name == "self":
continue
if param.kind in {inspect.Parameter.VAR_POSITIONAL, inspect.Parameter.VAR_KEYWORD}:
continue
default = param.default if param.default is not inspect._empty else None
annotation = param.annotation if param.annotation is not inspect._empty else None
if param.name in ("axis",):
continue
info: Dict[str, Any] = {
"name": param.name,
"default": default,
"annotation": annotation,
"description": "",
"group": "",
"nullable": False,
"adjusted": False,
"adjust_reason": "",
}
# YAML tunable choices
if param.name in tunable_ranges:
choices = tunable_ranges[param.name]
if isinstance(choices, list) and len(choices) > 1:
if all(isinstance(c, str) for c in choices):
info["type"] = "choice"
info["choices"] = choices
if default not in choices:
choices.insert(0, default) if default is not None else None
elif all(isinstance(c, (int, float)) for c in choices):
if isinstance(default, bool):
info["type"] = "bool"
elif isinstance(default, int) and all(isinstance(c, int) for c in choices):
info["type"] = "int"
info["min"] = min(choices)
info["max"] = max(choices)
info["step"] = 1
info["choices"] = choices
else:
info["type"] = "float"
info["min"] = float(min(choices))
info["max"] = float(max(choices))
info["step"] = round((max(choices) - min(choices)) / 20, 6)
info["choices"] = choices
else:
info["type"] = "choice"
info["choices"] = [str(c) for c in choices]
params.append(info)
continue
# Fallback: infer from default
if isinstance(default, bool):
info["type"] = "bool"
elif isinstance(default, int):
info["type"] = "int"
lower = 0 if default >= 0 else int(default * 3)
upper = default * 3 + 10 if default >= 0 else -default * 3 + 10
upper = max(upper, lower + 5)
info["min"] = lower
info["max"] = upper
info["step"] = 1
elif isinstance(default, float):
info["type"] = "float"
span = max(abs(default), 0.1)
info["min"] = round(float(max(0, default - 3 * span)), 6)
info["max"] = round(float(default + 3 * span), 6)
info["step"] = round(span / 20, 6)
elif isinstance(default, str):
info["type"] = "str"
elif default is None:
info["type"] = "text"
elif isinstance(default, (list, tuple, dict)):
info["type"] = "text"
else:
info["type"] = "text"
params.append(info)
return params
# ---------------------------------------------------------------------------
# Validation & adjustment
# ---------------------------------------------------------------------------
def validate_params_for_detector(
cls: type,
params: Dict[str, Any],
data_ctx: Dict[str, Any] | None = None,
) -> Tuple[Dict[str, Any], List[str], List[str]]:
"""Validate and auto-adjust parameters before running a detector.
Parameters
----------
cls : type
Detector class.
params : dict
User-supplied parameter values.
data_ctx : dict, optional
Data context for data-dependent constraints.
Returns
-------
adjusted_params : dict
Parameters with out-of-bounds values clamped.
warnings : list[str]
Human-readable warnings about auto-adjusted values.
errors : list[str]
Hard errors that cannot be auto-fixed.
"""
if not _has_parameter_schema(cls):
return params, [], []
from tsseg.algorithms.param_schema import get_ui_hints
from tsseg.algorithms.param_schema import validate_params as _validate
hints = get_ui_hints(cls, data_ctx=data_ctx)
adjusted = dict(params)
warnings: List[str] = []
errors: List[str] = []
# --- Auto-clamp numeric values to effective bounds --------------------
for pname, h in hints.items():
if pname not in adjusted:
continue
val = adjusted[pname]
if val is None:
continue
lo = h.get("min")
hi = h.get("max")
ptype = h.get("type", "any")
if ptype in ("int", "float") and (lo is not None or hi is not None):
clamped, was = _clamp_default(val, lo, hi, ptype)
if was:
adjusted[pname] = clamped
if data_ctx:
warnings.append(
f"**{pname}** = {val} → {clamped} (series: {data_ctx.get('n_samples', '?')} points)"
)
else:
warnings.append(f"**{pname}** = {val} → {clamped} (out of bounds)")
# --- Run the full schema validator for cross-constraints ---------------
try:
# Build a temporary instance to use validate_params
instance = cls.__new__(cls)
# Set params so get_params works (sklearn interface)
for k, v in adjusted.items():
setattr(instance, k, v)
errs = _validate(instance, data_ctx=data_ctx)
for e in errs:
# If the error is about a value we already clamped, downgrade to warning
errors.append(e)
except Exception as exc:
logger.debug(f"Schema validation skipped: {exc}")
return adjusted, warnings, errors