Skip to content

MultiSeedTrainerView

Multi-seed evaluation/plotting view for a Trainer.

Obtain via trainer.multi_seed(...). Exposes the same eval/plot methods as the underlying trainer, running each across convergent seed checkpoints and aggregating.

Source code in gradiend/trainer/core/multi_seed.py
def __init__(
    self,
    trainer: Any,
    *,
    selection: str = "all_convergent",
    aggregate: str = "mean",
    dispersion: str = "std",
    return_per_seed: bool = False,
) -> None:
    _validate_aggregate_dispersion(aggregate, dispersion)
    self._trainer = trainer
    self.selection = str(selection).strip().lower()
    self.aggregate = aggregate
    self.dispersion = dispersion
    self.return_per_seed = bool(return_per_seed)
    self._entries = resolve_seed_run_entries(trainer, self.selection)
    self._shared_base_model: Any = None
    self._shared_tokenizer: Any = None

_entries instance-attribute

_entries = resolve_seed_run_entries(trainer, self.selection)

_shared_base_model instance-attribute

_shared_base_model = None

_shared_tokenizer instance-attribute

_shared_tokenizer = None

_trainer instance-attribute

_trainer = trainer

aggregate instance-attribute

aggregate = aggregate

dispersion instance-attribute

dispersion = dispersion

return_per_seed instance-attribute

return_per_seed = bool(return_per_seed)

selection instance-attribute

selection = str(selection).strip().lower()

trainer property

trainer

__getattr__

__getattr__(name)
Source code in gradiend/trainer/core/multi_seed.py
def __getattr__(self, name: str) -> Any:
    if name.startswith("_"):
        raise AttributeError(name)
    if name in MULTI_SEED_CAPABLE_METHODS:
        return self._bind_method(name)
    return getattr(self._trainer, name)

_bind_method

_bind_method(name)
Source code in gradiend/trainer/core/multi_seed.py
def _bind_method(self, name: str) -> Callable[..., Any]:
    fn = getattr(self._trainer, name)

    def _caller(**kwargs: Any) -> Any:
        return self._run_for_seeds(name, fn, **kwargs)

    _caller.__name__ = name
    _caller.__doc__ = fn.__doc__
    return _caller

_cleanup_model

_cleanup_model(model)
Source code in gradiend/trainer/core/multi_seed.py
def _cleanup_model(self, model: Any) -> None:
    del model
    gc.collect()
    if torch.cuda.is_available():
        torch.cuda.empty_cache()

_load_seed_model

_load_seed_model(seed_path)
Source code in gradiend/trainer/core/multi_seed.py
def _load_seed_model(self, seed_path: str) -> Any:
    load_kwargs: Dict[str, Any] = {"use_cache": False}
    if self._shared_base_model is not None:
        load_kwargs["base_model"] = self._shared_base_model
    if self._shared_tokenizer is not None:
        load_kwargs["tokenizer"] = self._shared_tokenizer
    model = self._trainer.load_model(seed_path, **load_kwargs)
    if self._shared_base_model is None and getattr(model, "base_model", None) is not None:
        self._shared_base_model = model.base_model
    if self._shared_tokenizer is None and getattr(model, "tokenizer", None) is not None:
        self._shared_tokenizer = model.tokenizer
    return model

_plot_with_seed_encoder_df

_plot_with_seed_encoder_df(method_name, **kwargs)
Source code in gradiend/trainer/core/multi_seed.py
def _plot_with_seed_encoder_df(self, method_name: str, **kwargs: Any) -> Dict[str, Any]:
    return_per_seed = self.return_per_seed
    if "return_per_seed" in kwargs:
        return_per_seed = bool(kwargs.pop("return_per_seed"))
    eval_kwargs = {
        key: kwargs.pop(key)
        for key in list(kwargs)
        if key in ENCODER_DF_PLOT_EVAL_KEYS
    }
    eval_kwargs.setdefault("split", "test")
    eval_kwargs.setdefault("max_size", None)
    eval_kwargs.setdefault("use_cache", True)
    eval_kwargs["return_df"] = True
    eval_kwargs["plot"] = False

    results: List[Any] = []
    seed_values: List[int] = []
    for seed_val, seed_path in self._entries:
        model = self._load_seed_model(seed_path)
        try:
            self._prepare_seed_data(seed_val)
            with _seed_execution_context(self._trainer, seed_path, model):
                eval_result = self._trainer.evaluate_encoder(**eval_kwargs)
                encoder_df = eval_result.get("encoder_df") if isinstance(eval_result, dict) else None
                result = getattr(self._trainer, method_name)(encoder_df=encoder_df, **kwargs)
            results.append(result)
            seed_values.append(int(seed_val))
        finally:
            self._cleanup_model(model)

    payload = _aggregate_plot_results(
        results,
        seed_values,
        selection=self.selection,
        aggregate=self.aggregate,
        dispersion=self.dispersion,
        return_per_seed=return_per_seed,
    )
    if not return_per_seed:
        payload["seeds"].pop("per_seed", None)
    return payload

_prepare_seed_data

_prepare_seed_data(seed_value)
Source code in gradiend/trainer/core/multi_seed.py
def _prepare_seed_data(self, seed_value: int) -> None:
    refresh_splits = getattr(self._trainer, "_refresh_data_splits_for_seed", None)
    args = getattr(self._trainer, "_training_args", None) or getattr(self._trainer, "training_args", None)
    if callable(refresh_splits) and args is not None:
        metadata = self._seed_run_metadata(seed_value)
        refresh_splits(
            int(seed_value),
            args,
            split_cycle_index=metadata.get("split_cycle_index"),
            split_cycle_length=metadata.get("split_cycle_length"),
        )

_run_for_seeds

_run_for_seeds(method_name, fn, /, **kwargs)
Source code in gradiend/trainer/core/multi_seed.py
def _run_for_seeds(self, method_name: str, fn: Callable[..., Any], /, **kwargs: Any) -> Any:
    return_per_seed = self.return_per_seed
    if "return_per_seed" in kwargs:
        return_per_seed = bool(kwargs.pop("return_per_seed"))

    if len(self._entries) == 1:
        seed_val, seed_path = self._entries[0]
        model = self._load_seed_model(seed_path)
        try:
            self._prepare_seed_data(seed_val)
            with _seed_execution_context(self._trainer, seed_path, model):
                single = fn(**kwargs)
        finally:
            self._cleanup_model(model)
        if method_name in PLOT_METHODS:
            plot_payload = _aggregate_plot_results(
                [single],
                [seed_val],
                selection=self.selection,
                aggregate=self.aggregate,
                dispersion=self.dispersion,
                return_per_seed=return_per_seed,
            )
            if not return_per_seed:
                plot_payload["seeds"].pop("per_seed", None)
            return plot_payload
        if isinstance(single, dict):
            payload = dict(single)
            per_seed_payload = dict(payload)
            payload["seeds"] = {
                "n": 1,
                "values": [seed_val],
                "selection": self.selection,
                "aggregate": self.aggregate,
                "dispersion": self.dispersion,
                "stats": _build_stats_from_single(payload, self.aggregate, self.dispersion),
            }
            if return_per_seed:
                payload["seeds"]["per_seed"] = {seed_val: per_seed_payload}
            return payload
        return single

    results: List[Any] = []
    seed_values: List[int] = []
    for seed_val, seed_path in self._entries:
        model = self._load_seed_model(seed_path)
        try:
            self._prepare_seed_data(seed_val)
            with _seed_execution_context(self._trainer, seed_path, model):
                results.append(fn(**kwargs))
            seed_values.append(int(seed_val))
        finally:
            self._cleanup_model(model)

    if method_name in PLOT_METHODS:
        payload = _aggregate_plot_results(
            results,
            seed_values,
            selection=self.selection,
            aggregate=self.aggregate,
            dispersion=self.dispersion,
            return_per_seed=return_per_seed,
        )
        if not return_per_seed:
            payload["seeds"].pop("per_seed", None)
        return payload

    dict_results = [r for r in results if isinstance(r, dict)]
    if len(dict_results) != len(results):
        return {
            "results": results,
            "seeds": {
                "n": len(results),
                "values": seed_values,
                "selection": self.selection,
                "aggregate": self.aggregate,
                "dispersion": self.dispersion,
                "per_seed": dict(zip(seed_values, results)) if return_per_seed else None,
            },
        }

    if method_name == "evaluate":
        enc_results = [r.get("encoder", {}) for r in dict_results if isinstance(r.get("encoder"), dict)]
        dec_results = [r.get("decoder", {}) for r in dict_results if isinstance(r.get("decoder"), dict)]
        merged: Dict[str, Any] = {}
        if enc_results:
            merged["encoder"] = aggregate_eval_results(
                enc_results,
                seed_values,
                selection=self.selection,
                aggregate=self.aggregate,
                dispersion=self.dispersion,
                return_per_seed=return_per_seed,
            )
        if dec_results:
            merged["decoder"] = aggregate_eval_results(
                dec_results,
                seed_values,
                selection=self.selection,
                aggregate=self.aggregate,
                dispersion=self.dispersion,
                return_per_seed=return_per_seed,
            )
        return merged

    merged_eval = aggregate_eval_results(
        dict_results,
        seed_values,
        selection=self.selection,
        aggregate=self.aggregate,
        dispersion=self.dispersion,
        return_per_seed=return_per_seed,
    )
    if not return_per_seed and "per_seed" in merged_eval.get("seeds", {}):
        merged_eval["seeds"].pop("per_seed", None)
    return merged_eval

_seed_run_metadata

_seed_run_metadata(seed_value)
Source code in gradiend/trainer/core/multi_seed.py
def _seed_run_metadata(self, seed_value: int) -> Dict[str, Any]:
    if not hasattr(self._trainer, "get_seed_report"):
        return {}
    report = self._trainer.get_seed_report()
    runs = report.get("runs", []) if isinstance(report, dict) else []
    if not isinstance(runs, list):
        return {}
    for run in runs:
        if isinstance(run, dict) and run.get("seed") == int(seed_value):
            return run
    return {}

evaluate

evaluate(**kwargs)

Run trainer.evaluate for each selected seed and aggregate results.

Parameters:

Name Type Description Default
**kwargs Any

Forwarded to the underlying trainer method.

{}
Source code in gradiend/trainer/core/multi_seed.py
def evaluate(self, **kwargs: Any) -> Dict[str, Any]:
    """Run ``trainer.evaluate`` for each selected seed and aggregate results.

    Args:
        **kwargs: Forwarded to the underlying trainer method.
    """
    return self._bind_method("evaluate")(**kwargs)

evaluate_decoder

evaluate_decoder(**kwargs)

Run decoder evaluation for each selected seed and aggregate metrics.

Parameters:

Name Type Description Default
**kwargs Any

Forwarded to trainer.evaluate_decoder.

{}
Source code in gradiend/trainer/core/multi_seed.py
def evaluate_decoder(self, **kwargs: Any) -> Dict[str, Any]:
    """Run decoder evaluation for each selected seed and aggregate metrics.

    Args:
        **kwargs: Forwarded to ``trainer.evaluate_decoder``.
    """
    return self._bind_method("evaluate_decoder")(**kwargs)

evaluate_encoder

evaluate_encoder(**kwargs)

Run encoder evaluation for each selected seed and aggregate metrics.

Parameters:

Name Type Description Default
**kwargs Any

Forwarded to trainer.evaluate_encoder.

{}
Source code in gradiend/trainer/core/multi_seed.py
def evaluate_encoder(self, **kwargs: Any) -> Dict[str, Any]:
    """Run encoder evaluation for each selected seed and aggregate metrics.

    Args:
        **kwargs: Forwarded to ``trainer.evaluate_encoder``.
    """
    return self._bind_method("evaluate_encoder")(**kwargs)

get_model

get_model(use_cache=None, *, load_directory=None, gradiend_only=False, **kwargs)

Load the selected seed checkpoint(s) for analysis or comparison.

A single checkpoint returns the model directly. Multiple checkpoints return a :class:~gradiend.trainer.core.seed_models.SeedModelGroup for use with similarity / top-k overlap matrix builders.

Source code in gradiend/trainer/core/multi_seed.py
def get_model(
    self,
    use_cache: Optional[bool] = None,
    *,
    load_directory: Optional[Any] = None,
    gradiend_only: bool = False,
    **kwargs: Any,
) -> Any:
    """Load the selected seed checkpoint(s) for analysis or comparison.

    A single checkpoint returns the model directly. Multiple checkpoints
    return a :class:`~gradiend.trainer.core.seed_models.SeedModelGroup` for
    use with similarity / top-k overlap matrix builders.
    """
    from gradiend.trainer.core.seed_models import SeedModelGroup

    if load_directory is not None:
        return self._trainer.get_model(
            use_cache=use_cache,
            load_directory=load_directory,
            **kwargs,
        )
    if len(self._entries) <= 1:
        seed_path = self._entries[0][1]
        if gradiend_only:
            from gradiend.trainer.suite.definitions import _load_gradiend_only_model

            return _load_gradiend_only_model(seed_path, device="cpu")
        return self._load_seed_model(seed_path)

    if gradiend_only:
        from gradiend.trainer.suite.definitions import _load_gradiend_only_model

        models = [
            _load_gradiend_only_model(path, device="cpu")
            for _, path in self._entries
        ]
    else:
        models = self.load_models()
    return SeedModelGroup(
        models,
        selection=self.selection,
        aggregate=self.aggregate,
        dispersion=self.dispersion,
        seed_values=self.seed_values(),
    )

load_models

load_models()

Load all seed checkpoints into memory (caller should release when done).

Source code in gradiend/trainer/core/multi_seed.py
def load_models(self) -> List[Any]:
    """Load all seed checkpoints into memory (caller should release when done)."""
    models, _, _ = load_seed_model_group(
        self._trainer,
        selection=self.selection,
        shared_base_model=self._shared_base_model,
        shared_tokenizer=self._shared_tokenizer,
    )
    if models and self._shared_base_model is None:
        first = models[0]
        if getattr(first, "base_model", None) is not None:
            self._shared_base_model = first.base_model
        if getattr(first, "tokenizer", None) is not None:
            self._shared_tokenizer = first.tokenizer
    return models

multi_seed

multi_seed(**kwargs)

Re-wrap the inner trainer (replaces view options when kwargs are passed).

Source code in gradiend/trainer/core/multi_seed.py
def multi_seed(self, **kwargs: Any) -> "MultiSeedTrainerView":
    """Re-wrap the inner trainer (replaces view options when kwargs are passed)."""
    if kwargs:
        return self._trainer.multi_seed(**kwargs)
    return self

plot_encoder_by_target

plot_encoder_by_target(**kwargs)

Create one held-out-target encoder plot with one row per selected seed.

The default selection is "all_convergent", so this method naturally focuses the target-word analysis on convergent checkpoints.

Source code in gradiend/trainer/core/multi_seed.py
def plot_encoder_by_target(self, **kwargs: Any) -> Dict[str, Any]:
    """Create one held-out-target encoder plot with one row per selected seed.

    The default selection is ``"all_convergent"``, so this method naturally
    focuses the target-word analysis on convergent checkpoints.
    """
    import pandas as pd

    from gradiend.visualizer.encoder_by_target import plot_encoder_by_target_seed_grid

    plot_kwargs = {
        key: kwargs.pop(key)
        for key in list(kwargs)
        if key
        in {
            "target_col",
            "class_col",
            "hue_col",
            "class_order",
            "output",
            "output_dir",
            "experiment_dir",
            "show",
            "figsize",
            "jitter",
            "point_size",
            "title",
            "plot_style",
            "interactive",
            "height",
            "error_stat",
            "show_seed_points",
            "error_group_by_split",
            "combine_seed_rows",
        }
    }
    eval_kwargs = dict(kwargs)
    eval_kwargs.setdefault("split", "all")
    eval_kwargs.setdefault("max_size", None)
    eval_kwargs.setdefault("use_cache", True)
    eval_kwargs["return_df"] = True
    eval_kwargs["plot"] = False

    frames: List[Any] = []
    seed_values: List[int] = []
    for seed_val, seed_path in self._entries:
        model = self._load_seed_model(seed_path)
        try:
            self._prepare_seed_data(seed_val)
            with _seed_execution_context(self._trainer, seed_path, model):
                result = self._trainer.evaluate_encoder(**eval_kwargs)
            frame = result.get("encoder_df") if isinstance(result, dict) else None
            if frame is not None and not frame.empty:
                frame = frame.copy()
                frame["seed"] = int(seed_val)
                frames.append(frame)
                seed_values.append(int(seed_val))
        finally:
            self._cleanup_model(model)

    if not frames:
        return {
            "path": None,
            "seeds": {
                "n": 0,
                "values": [],
                "selection": self.selection,
                "aggregate": self.aggregate,
                "dispersion": self.dispersion,
            },
        }

    encoder_df = pd.concat(frames, ignore_index=True)
    id2label = dict(getattr(self._trainer, "_id2label", None) or {})
    config_obj = getattr(self._trainer, "config", None)
    config_map = getattr(config_obj, "id2label", None) if config_obj is not None else None
    if isinstance(config_map, dict):
        id2label.update(config_map)
    if "class_order" not in plot_kwargs:
        pair = getattr(self._trainer, "pair", None)
        if pair:
            plot_kwargs["class_order"] = list(pair)
    path = plot_encoder_by_target_seed_grid(
        encoder_df,
        id2label=id2label or None,
        **plot_kwargs,
    )
    return {
        "path": path,
        "seeds": {
            "n": len(seed_values),
            "values": seed_values,
            "selection": self.selection,
            "aggregate": self.aggregate,
            "dispersion": self.dispersion,
        },
    }

plot_encoder_distributions

plot_encoder_distributions(**kwargs)

Create encoder-distribution plots for each selected seed.

Parameters:

Name Type Description Default
**kwargs Any

Forwarded to trainer.plot_encoder_distributions.

{}
Source code in gradiend/trainer/core/multi_seed.py
def plot_encoder_distributions(self, **kwargs: Any) -> Dict[str, Any]:
    """Create encoder-distribution plots for each selected seed.

    Args:
        **kwargs: Forwarded to ``trainer.plot_encoder_distributions``.
    """
    return self._plot_with_seed_encoder_df("plot_encoder_distributions", **kwargs)

plot_encoder_scatter

plot_encoder_scatter(**kwargs)

Create interactive encoder-scatter plots for each selected seed.

Parameters:

Name Type Description Default
**kwargs Any

Forwarded to trainer.plot_encoder_scatter.

{}
Source code in gradiend/trainer/core/multi_seed.py
def plot_encoder_scatter(self, **kwargs: Any) -> Dict[str, Any]:
    """Create interactive encoder-scatter plots for each selected seed.

    Args:
        **kwargs: Forwarded to ``trainer.plot_encoder_scatter``.
    """
    return self._plot_with_seed_encoder_df("plot_encoder_scatter", **kwargs)

plot_encoder_strip_by_split

plot_encoder_strip_by_split(**kwargs)

Create encoder strip-by-split plots for each selected seed.

Parameters:

Name Type Description Default
**kwargs Any

Forwarded to trainer.plot_encoder_strip_by_split.

{}
Source code in gradiend/trainer/core/multi_seed.py
def plot_encoder_strip_by_split(self, **kwargs: Any) -> Dict[str, Any]:
    """Create encoder strip-by-split plots for each selected seed.

    Args:
        **kwargs: Forwarded to ``trainer.plot_encoder_strip_by_split``.
    """
    return self._plot_with_seed_encoder_df("plot_encoder_strip_by_split", **kwargs)

plot_probability_shifts

plot_probability_shifts(**kwargs)

Create decoder probability-shift plots for each selected seed.

Parameters:

Name Type Description Default
**kwargs Any

Forwarded to trainer.plot_probability_shifts.

{}
Source code in gradiend/trainer/core/multi_seed.py
def plot_probability_shifts(self, **kwargs: Any) -> Dict[str, Any]:
    """Create decoder probability-shift plots for each selected seed.

    Args:
        **kwargs: Forwarded to ``trainer.plot_probability_shifts``.
    """
    return self._bind_method("plot_probability_shifts")(**kwargs)

plot_training_convergence

plot_training_convergence(**kwargs)

Create training-convergence plots for each selected seed.

Parameters:

Name Type Description Default
**kwargs Any

Forwarded to trainer.plot_training_convergence.

{}
Source code in gradiend/trainer/core/multi_seed.py
def plot_training_convergence(self, **kwargs: Any) -> Dict[str, Any]:
    """Create training-convergence plots for each selected seed.

    Args:
        **kwargs: Forwarded to ``trainer.plot_training_convergence``.
    """
    return self._bind_method("plot_training_convergence")(**kwargs)

seed_models

seed_models()

Lazy-load seed checkpoints (shared base model when possible).

Source code in gradiend/trainer/core/multi_seed.py
def seed_models(self) -> Iterator[Any]:
    """Lazy-load seed checkpoints (shared base model when possible)."""
    for _seed_val, seed_path in self._entries:
        model = self._load_seed_model(seed_path)
        try:
            yield model
        finally:
            self._cleanup_model(model)

seed_paths

seed_paths()
Source code in gradiend/trainer/core/multi_seed.py
def seed_paths(self) -> List[str]:
    return [path for _, path in self._entries]

seed_values

seed_values()
Source code in gradiend/trainer/core/multi_seed.py
def seed_values(self) -> List[int]:
    return [seed for seed, _ in self._entries]