Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 9 additions & 1 deletion docs/user-guide/running.md
Original file line number Diff line number Diff line change
Expand Up @@ -255,7 +255,15 @@ safe-synthesizer run generate \
| `--run-path` | Explicit path to a previous run's output directory |
| `--wandb-resume-job-id` | WandB run ID to resume (or path to file containing the ID) |

Accepts the same common options and synthesis parameter overrides as `run`.
Accepts the same common options and synthesis parameter override syntax as `run`.

!!! note "Override scope on resume"
`run generate` reloads the trained run's saved configuration. Only
`generation` and `evaluation` overrides (and `emit_telemetry`) supplied via
`--config` or `--section__field` flags take effect; fields you do not set
keep their saved values. `training`, `data`, `privacy`, and `time_series`
are always inherited from the trained run and cannot be changed at generate
time, since they describe how the adapter was produced.

### `run --validate`

Expand Down
2 changes: 1 addition & 1 deletion src/nemo_safe_synthesizer/cli/run.py
Original file line number Diff line number Diff line change
Expand Up @@ -572,7 +572,7 @@ def run_generate(

try:
nss = (
nss.load_from_save_path()
nss.load_from_save_path(runtime_config=config)
.process_data()
.generate()
.evaluate()
Expand Down
8 changes: 4 additions & 4 deletions src/nemo_safe_synthesizer/cli/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -276,7 +276,8 @@ def common_setup(
df: pd.DataFrame | None = None
if settings.data_source:
dataset_info = dataset_registry.get_dataset(settings.data_source)
synthesis_overrides = merge_dicts(synthesis_overrides, dataset_info.overrides or dict())
if not resume:
synthesis_overrides = merge_dicts(synthesis_overrides, dataset_info.overrides or dict())
df = dataset_info.fetch()
elif resume:
# For generate-only runs without --data-source, verify cached dataset exists.
Expand Down Expand Up @@ -447,9 +448,8 @@ def merge_overrides(config_path: str | Path | None, overrides: dict) -> SafeSynt
if config_path is None:
my_config = SafeSynthesizerParameters.model_validate(overrides)
else:
params = merge_dicts(
SafeSynthesizerParameters.from_yaml(config_path).model_dump(exclude_unset=False), overrides
)
file_config = SafeSynthesizerParameters.from_yaml(config_path).model_dump(exclude_unset=True)
params = merge_dicts(file_config, overrides)
my_config = SafeSynthesizerParameters.model_validate(params)
except ValidationError as e:
click.echo(f"{config_path} is invalid:\n{e}")
Expand Down
74 changes: 73 additions & 1 deletion src/nemo_safe_synthesizer/config/parameters.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,12 +6,13 @@
import warnings
from typing import Any, Self

from pydantic import Field, model_validator
from pydantic import BaseModel, Field, model_validator

from ..configurator.parameters import Parameters
from ..errors import ParameterError
from ..observability import get_logger
from ..telemetry import _telemetry_enabled
from ..utils import merge_dicts
from .data import DataParameters
from .differential_privacy import DifferentialPrivacyHyperparams
from .evaluate import EvaluationParameters
Expand All @@ -28,6 +29,42 @@
logger = get_logger(__name__)


def _collect_set_fields(model: BaseModel) -> dict[str, Any]:
"""Recursively collect a model's explicitly-set fields as a nested dict.

Unlike ``model_dump(exclude_unset=True)``, nested models are always
traversed -- even when the parent did not mark the nested field as set --
so in-place mutations of nested fields (e.g.
``cfg.generation.validation.foo = True``) are captured. A nested model is
included only when it has at least one set field of its own.
"""
overrides: dict[str, Any] = {}
for name in type(model).model_fields:
value = getattr(model, name)
if isinstance(value, BaseModel):
nested = _collect_set_fields(value)
if nested:
overrides[name] = nested
elif name in model.__pydantic_fields_set__:
overrides[name] = value
return overrides


def _overlay_set_fields(saved: Parameters, runtime: Parameters) -> Parameters:
"""Deep-merge ``runtime``'s explicitly-set fields onto ``saved``.

Only fields ``runtime`` marks as set (recursively, at every nesting level)
override ``saved``; unset fields keep their saved values. The merged mapping
is revalidated through the model so nested groups and type coercion are
handled correctly. Returns ``saved`` unchanged when ``runtime`` sets no
fields.
"""
overrides = _collect_set_fields(runtime)
if not overrides:
return saved
return saved.model_validate(merge_dicts(saved.model_dump(), overrides))


class SafeSynthesizerParameters(Parameters):
"""Main configuration class for the Safe Synthesizer pipeline.

Expand Down Expand Up @@ -199,3 +236,38 @@ def from_params(cls, **kwargs) -> "SafeSynthesizerParameters":
if "emit_telemetry" in kwargs:
extra["emit_telemetry"] = kwargs["emit_telemetry"]
return cls(**extra)

def with_runtime_overrides(self, runtime: SafeSynthesizerParameters) -> "SafeSynthesizerParameters":
"""Apply resume-time generation/evaluation/telemetry overrides onto a copy of self.

``self`` is the saved training-run config. Only explicitly-set
``generation`` and ``evaluation`` fields from ``runtime`` are merged in,
plus ``emit_telemetry`` when the caller set it. Training, data, privacy,
and other sections are preserved so training provenance survives a
generate-only resume.

Args:
runtime: Config carrying resume-time CLI/SDK overrides. Typically
sparse -- only the fields the caller set are applied.

Returns:
A new ``SafeSynthesizerParameters`` with overrides applied. The
result is fully independent of ``self``: sections that are not
overridden are deep-copied, so later mutation of either object does
not affect the other.
"""
updates: dict[str, object] = {}
# Only record sections that actually changed; unchanged sections are
# deep-copied by ``model_copy(deep=True)`` below so the returned config
# never shares mutable sub-objects with ``self``.
generation = _overlay_set_fields(self.generation, runtime.generation)
if generation is not self.generation:
updates["generation"] = generation
evaluation = _overlay_set_fields(self.evaluation, runtime.evaluation)
if evaluation is not self.evaluation:
updates["evaluation"] = evaluation
# emit_telemetry is a top-level scalar: detect explicit assignment,
# since there is no sub-model to inspect for set fields.
if "emit_telemetry" in runtime.__pydantic_fields_set__:
updates["emit_telemetry"] = runtime.emit_telemetry
return self.model_copy(update=updates, deep=True)
47 changes: 45 additions & 2 deletions src/nemo_safe_synthesizer/sdk/library_builder.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@

import os
import time
from collections.abc import Iterator
from pathlib import Path
from typing import TYPE_CHECKING

Expand All @@ -18,6 +19,7 @@
SafeSynthesizerParameters,
)
from ..config.autoconfig import AutoConfigResolver
from ..configurator.parameters import Parameters
from ..errors import ParameterError
from ..evaluation.evaluator import Evaluator
from ..generation.timeseries_backend import TimeseriesBackend
Expand Down Expand Up @@ -51,6 +53,37 @@
from ..training.backend import TrainingBackend


def _direct_parameter_values(params: Parameters) -> Iterator[tuple[str, object]]:
for item in params._iter_parameters(recursive=False):
yield next(iter(item.items()))


def _default_drift_paths(saved_model: Parameters, default_model: Parameters, prefix: str = "") -> list[str]:
paths: list[str] = []
default_values = dict(_direct_parameter_values(default_model))
for field_name, saved_value in _direct_parameter_values(saved_model):
default_value = default_values[field_name]
field_path = f"{prefix}.{field_name}" if prefix else field_name

saved_attr = saved_model.__dict__[field_name]
default_attr = default_model.__dict__[field_name]
if isinstance(saved_attr, Parameters) and isinstance(default_attr, Parameters):
paths.extend(_default_drift_paths(saved_attr, default_attr, field_path))
elif saved_value != default_value:
paths.append(field_path)
return paths


def _warn_for_saved_default_drift(saved_config: SafeSynthesizerParameters) -> None:
"""Warn when saved materialized values differ from current defaults."""
current_defaults = SafeSynthesizerParameters()
for path in _default_drift_paths(saved_config, current_defaults):
logger.user.warning(
f"Saved run config value at {path} is non-default; preserving saved value.",
extra={"config_path": path},
)
Comment thread
coderabbitai[bot] marked this conversation as resolved.


def _build_telemetry_event(ss: SafeSynthesizer, status: TaskStatusEnum) -> NSSTrainingAndGenerationEvent:
"""Build a telemetry event from the current pipeline state."""
cfg = ss._nss_config
Expand Down Expand Up @@ -237,12 +270,15 @@ def _ensure_observability(self) -> None:
initialize_observability()

@traced("SafeSynthesizer.load_from_save_path", category=LogCategory.RUNTIME)
def load_from_save_path(self) -> SafeSynthesizer:
def load_from_save_path(self, runtime_config: SafeSynthesizerParameters | None = None) -> SafeSynthesizer:
"""Load the Safe Synthesizer configuration from the save path.

Loads the configuration from the source run directory's config file.
When resuming from a trained model for generation, the source paths
point to the parent workdir that contains the trained adapter.
Optional ``runtime_config`` values for generation and evaluation are
applied after loading the saved training-run config so resume-time CLI
overrides work without mutating the persisted train config.

Always prefers cached train/test splits from the training run to ensure
evaluation metrics are consistent and privacy guarantees are maintained.
Expand All @@ -256,7 +292,14 @@ def load_from_save_path(self) -> SafeSynthesizer:
# Use source paths which point to parent workdir when resuming for generation
config_file = self._workdir.source_config

self._nss_config = SafeSynthesizerParameters.from_json(config_file)
saved_config = SafeSynthesizerParameters.from_json(config_file)
_warn_for_saved_default_drift(saved_config)
if runtime_config is not None:
saved_config = saved_config.with_runtime_overrides(runtime_config)
self._nss_config = saved_config
self._generation_config = self._nss_config.generation
self._evaluation_config = self._nss_config.evaluation
self._emit_telemetry_config = self._nss_config.emit_telemetry

# Load model metadata from saved file (contains initial_prefill for timeseries)
# rather than creating new metadata from config
Expand Down
2 changes: 2 additions & 0 deletions tests/cli/test_run.py
Original file line number Diff line number Diff line change
Expand Up @@ -633,6 +633,7 @@ def test_generate_with_run_path_calls_common_setup(
dummy_csv: Path,
tmp_path: Path,
patched_run_dependencies: dict,
mock_common_setup_return: tuple,
):
"""Verify generate with --run-path calls common_setup correctly."""
run_dir = tmp_path / "trained-run"
Expand All @@ -659,6 +660,7 @@ def test_generate_with_run_path_calls_common_setup(
settings: CLISettings = call_kwargs["settings"]
assert settings.run_path == str(run_dir)
mock_ss = patched_run_dependencies["safe_synthesizer"]
mock_ss.load_from_save_path.assert_called_once_with(runtime_config=mock_common_setup_return[1])
mock_ss.evaluate.assert_called_once_with()
patched_run_dependencies["emit_telemetry"].assert_called_once_with(mock_ss, TaskStatusEnum.COMPLETED)

Expand Down
43 changes: 42 additions & 1 deletion tests/cli/test_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,7 @@
import pytest

from nemo_safe_synthesizer.cli.settings import CLISettings
from nemo_safe_synthesizer.cli.utils import _propagate_runtime_settings_to_env, common_setup
from nemo_safe_synthesizer.cli.utils import _propagate_runtime_settings_to_env, common_setup, merge_overrides


@pytest.fixture
Expand Down Expand Up @@ -230,6 +230,29 @@ def test_merge_dataset_and_cli_overrides(
assert config.training.batch_size == 8
assert config.data.holdout == 0.1

def test_resume_uses_registry_data_without_registry_config_overrides(
self,
registry_with_base_url: Path,
patched_common_setup_dependencies: dict,
):
"""Resume treats dataset registry overrides as default-like, not runtime overrides."""
settings = CLISettings.from_cli_kwargs(
data_source="test-data",
dataset_registry=str(registry_with_base_url),
synthesis_overrides={
"generation": {
"temperature": 0.7,
},
},
)

_, config, df, _ = common_setup(settings, resume=True)

assert df is not None
assert list(df.columns) == ["x", "y", "z"]
assert config.generation.num_records == 1000
assert config.generation.temperature == 0.7

def test_overrides_config_registry_and_cli(
self,
registry_with_dataset: Path,
Expand Down Expand Up @@ -327,6 +350,24 @@ def test_apply_cli_overrides_without_registry(
assert config.generation.temperature == 0.7


class TestMergeOverrides:
"""Tests for config-file and CLI override merging."""

def test_partial_config_preserves_explicit_field_metadata(self, tmp_path: Path):
"""Partial config files should not make default values look explicit."""
config_file = tmp_path / "config.yaml"
config_file.write_text("""
generation:
num_records: 77
""")

config = merge_overrides(config_file, {})

assert config.generation.num_records == 77
assert config.generation.use_structured_generation is False
assert config.model_dump(exclude_unset=True) == {"generation": {"num_records": 77}}


class TestPropagateRuntimeSettingsToEnv:
"""Tests for materializing CLISettings runtime fields back to os.environ."""

Expand Down
Loading
Loading