ezmsg.learn.collection.sample_adapt_regressor#
Module Attributes
Default torch model class used when |
Functions
- build_sample_adapt_regressor(settings)[source]#
Build a decode collection wired around a single regressor engine.
The regressor backend (River/sklearn, torch-mlp, or refit-Kalman) is selected from
settings.model_typeand the collection class is defined dynamically so the graph contains exactly the units that backend uses — no inert, declared-but-unwired units. The signal path (and, for the linear engine, the online-adaptation sample path) wire to that one unit, so there is no per- backend wiring to keep in sync.- Return type:
- Parameters:
settings (SampleAdaptRegressorSettings)
Classes
- class DecodeOutputAdapter(*args, settings=None, **kwargs)[source]#
Bases:
BaseTransformerUnit[DecodeOutputAdapterSettings,AxisArray,AxisArray,DecodeOutputAdapterProcessor]- Parameters:
settings (Settings | None)
- SETTINGS#
alias of
DecodeOutputAdapterSettings
- class DecodeOutputAdapterProcessor(*args, **kwargs)[source]#
Bases:
BaseStatefulTransformer[DecodeOutputAdapterSettings,AxisArray,AxisArray,DecodeOutputAdapterState]Normalize a decoder output into a
(time, ch)AxisArray.The torch (
{"output": ...}-keyed) and Kalman (["time", "state"]) engines emit differently-shaped outputs than the River/sklearn regressor. This rebuilds a uniform(time, ch=output_labels)message — keyed<input>_predlikeAdaptiveLinearRegressorUnit— so downstream consumers see one contract regardless of backend.
- class DecodeOutputAdapterSettings(output_labels=None)[source]#
Bases:
Settings- Parameters:
output_labels (list | None)
- class SampleAdaptRegressorSettings(model_type=AdaptiveLinearRegressor.LINEAR, model_path=None, model_kwargs=<factory>, model_class='ezmsg.learn.model.mlp.MLP', device=None, steady_state=True, output_labels=None, resample_axis='time', resample_buffer_duration=2.0, sampler_max_buffer_dur=5.0, decode_window_dur=None, decode_window_shift=None)[source]#
Bases:
Settings- Parameters:
model_type (AdaptiveLinearRegressor | str)
model_path (str | None)
model_kwargs (dict)
model_class (str)
device (str | None)
steady_state (bool)
output_labels (list | None)
resample_axis (str)
resample_buffer_duration (float)
sampler_max_buffer_dur (float)
decode_window_dur (float | None)
decode_window_shift (float | None)
- model_type: AdaptiveLinearRegressor | str = 'linear'#
Regressor backend/model.
- model_path: str | None = None#
Optional path to a pre-trained checkpoint. Format depends on the backend: a pickled River/sklearn estimator, a
torch.saveartifact (mlp), or a pickled state-space matrix dict (kalman).
- model_class: str = 'ezmsg.learn.model.mlp.MLP'#
Fully-qualified torch model class used when
model_type == "mlp".
- __init__(model_type=AdaptiveLinearRegressor.LINEAR, model_path=None, model_kwargs=<factory>, model_class='ezmsg.learn.model.mlp.MLP', device=None, steady_state=True, output_labels=None, resample_axis='time', resample_buffer_duration=2.0, sampler_max_buffer_dur=5.0, decode_window_dur=None, decode_window_shift=None)#
- Parameters:
model_type (AdaptiveLinearRegressor | str)
model_path (str | None)
model_kwargs (dict)
model_class (str)
device (str | None)
steady_state (bool)
output_labels (list | None)
resample_axis (str)
resample_buffer_duration (float)
sampler_max_buffer_dur (float)
decode_window_dur (float | None)
decode_window_shift (float | None)
- Return type:
None