yug 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
yug/__about__.py ADDED
@@ -0,0 +1,5 @@
1
+ # Copyright 2026 Birla AI Labs
2
+ # SPDX-License-Identifier: Apache-2.0
3
+ """Single source of truth for the package version."""
4
+
5
+ __version__ = "0.1.0"
yug/__init__.py ADDED
@@ -0,0 +1,57 @@
1
+ # Copyright 2026 Birla AI Labs
2
+ # SPDX-License-Identifier: Apache-2.0
3
+ """Yug — a foundation model for zero-shot probabilistic time-series forecasting.
4
+
5
+ from yug import YugPipeline
6
+
7
+ pipe = YugPipeline.from_pretrained("birlaailabs/yug")
8
+ forecast = pipe.predict(context, prediction_length=64)
9
+ forecast.median
10
+
11
+ Everything below :class:`YugPipeline` is available for the cases it does
12
+ not cover — :class:`YugForecaster` for raw per-anchor head output,
13
+ :class:`Yug_Model` for the network itself — but most users need only the
14
+ pipeline and the forecast it returns.
15
+ """
16
+
17
+ from .__about__ import __version__
18
+ from .base import BaseForecastPipeline, QuantileForecast
19
+ from .configs import InferenceConfig, PredictionConfig, YugConfig
20
+ from .engine import CachedUnivariateEngine
21
+ from .forecaster import ForecastResult, YugForecaster
22
+ from .model import Yug_Model
23
+ from .pipeline import YugPipeline
24
+ from .utils import QuantileSampler, freq_id, normalize_freq
25
+
26
+ #: Alias matching the class name used in the reference research scripts.
27
+ Yug = Yug_Model
28
+
29
+ __all__ = [
30
+ "__version__",
31
+ # Primary entry point
32
+ "YugPipeline",
33
+ "QuantileForecast",
34
+ "YugConfig",
35
+ # Configuration
36
+ "InferenceConfig",
37
+ "PredictionConfig",
38
+ # Lower-level building blocks
39
+ "BaseForecastPipeline",
40
+ "Yug_Model",
41
+ "Yug",
42
+ "YugForecaster",
43
+ "ForecastResult",
44
+ "CachedUnivariateEngine",
45
+ "QuantileSampler",
46
+ "normalize_freq",
47
+ "freq_id",
48
+ ]
49
+
50
+
51
+ def __getattr__(name: str):
52
+ """Expose the GluonTS adapter lazily so gluonts stays an optional extra."""
53
+ if name == "YugPredictor":
54
+ from .gluonts_adapter import YugPredictor
55
+
56
+ return YugPredictor
57
+ raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
yug/__main__.py ADDED
@@ -0,0 +1,162 @@
1
+ # Copyright 2026 Birla AI Labs
2
+ # SPDX-License-Identifier: Apache-2.0
3
+ """Pre-flight check: can this machine load and run Yug?
4
+
5
+ python -m yug # or: yug-check
6
+
7
+ Reports Python, PyTorch, accelerator, memory and disk, then says whether a
8
+ forecast will run and how fast to expect it to be. Exits non-zero if something
9
+ would actually prevent the model from loading, so it is usable in CI and by
10
+ agents as a gate before a long download.
11
+ """
12
+
13
+ from __future__ import annotations
14
+
15
+ import argparse
16
+ import platform
17
+ import shutil
18
+ import sys
19
+
20
+ # Weights are fp32; ~271.8M parameters is ~1.0 GB on disk and in memory, and
21
+ # activations plus the rollout KV cache need headroom on top.
22
+ WEIGHTS_GB = 1.15
23
+ MIN_RAM_GB = 3.0
24
+ RECOMMENDED_VRAM_GB = 4.0
25
+ MIN_DISK_GB = 2.0
26
+ MIN_PYTHON = (3, 10)
27
+
28
+
29
+ def _gb(n_bytes: float) -> float:
30
+ return n_bytes / (1024**3)
31
+
32
+
33
+ def _host_ram_gb() -> float | None:
34
+ """Total system RAM, or None when it cannot be determined portably."""
35
+ try:
36
+ return _gb(os_sysconf_ram())
37
+ except Exception:
38
+ return None
39
+
40
+
41
+ def os_sysconf_ram() -> float:
42
+ import os
43
+
44
+ return os.sysconf("SC_PAGE_SIZE") * os.sysconf("SC_PHYS_PAGES")
45
+
46
+
47
+ def main(argv: list[str] | None = None) -> int:
48
+ parser = argparse.ArgumentParser(
49
+ prog="yug-check",
50
+ description="Check whether this machine can load and run Yug.",
51
+ )
52
+ parser.add_argument(
53
+ "--quiet", action="store_true", help="print only the final verdict"
54
+ )
55
+ args = parser.parse_args(argv)
56
+
57
+ problems: list[str] = []
58
+ warnings: list[str] = []
59
+ lines: list[str] = []
60
+
61
+ def report(label: str, value: str) -> None:
62
+ lines.append(f" {label:<22} {value}")
63
+
64
+ lines.append("Yug pre-flight check")
65
+ lines.append("=" * 46)
66
+
67
+ # --- Python ------------------------------------------------------------
68
+ py = sys.version_info
69
+ report("Python", f"{py.major}.{py.minor}.{py.micro}")
70
+ report("Platform", f"{platform.system()} {platform.machine()}")
71
+ if (py.major, py.minor) < MIN_PYTHON:
72
+ problems.append(
73
+ f"Python {MIN_PYTHON[0]}.{MIN_PYTHON[1]}+ is required, found "
74
+ f"{py.major}.{py.minor}"
75
+ )
76
+
77
+ # --- PyTorch and accelerator -------------------------------------------
78
+ try:
79
+ import torch
80
+ except ImportError:
81
+ report("PyTorch", "NOT INSTALLED")
82
+ problems.append("PyTorch is not installed: pip install yug")
83
+ else:
84
+ report("PyTorch", torch.__version__)
85
+
86
+ if torch.cuda.is_available():
87
+ idx = torch.cuda.current_device()
88
+ props = torch.cuda.get_device_properties(idx)
89
+ vram = _gb(props.total_memory)
90
+ report("Accelerator", f"CUDA — {props.name}")
91
+ report("VRAM", f"{vram:.1f} GB")
92
+ report("CUDA runtime", torch.version.cuda or "unknown")
93
+ if vram < RECOMMENDED_VRAM_GB:
94
+ warnings.append(
95
+ f"{vram:.1f} GB of VRAM is below the {RECOMMENDED_VRAM_GB:.0f} GB "
96
+ f"recommended; reduce num_samples if you hit OOM"
97
+ )
98
+ elif getattr(torch.backends, "mps", None) and torch.backends.mps.is_available():
99
+ report("Accelerator", "Apple MPS")
100
+ warnings.append("running on MPS: supported, but noticeably slower than CUDA")
101
+ else:
102
+ report("Accelerator", "none — CPU only")
103
+ warnings.append(
104
+ "no accelerator found; a 512-step forecast takes minutes on CPU "
105
+ "rather than under a second on a GPU. Lower num_samples to compensate."
106
+ )
107
+
108
+ # --- Package -----------------------------------------------------------
109
+ try:
110
+ from yug import __version__
111
+
112
+ report("yug", __version__)
113
+ except ImportError:
114
+ report("yug", "NOT IMPORTABLE")
115
+ problems.append("yug cannot be imported: pip install yug")
116
+
117
+ # --- Memory and disk ---------------------------------------------------
118
+ ram = _host_ram_gb()
119
+ if ram is not None:
120
+ report("System RAM", f"{ram:.1f} GB")
121
+ if ram < MIN_RAM_GB:
122
+ problems.append(
123
+ f"{ram:.1f} GB of RAM is below the {MIN_RAM_GB:.0f} GB needed to "
124
+ f"hold the weights"
125
+ )
126
+
127
+ free_disk = _gb(shutil.disk_usage(".").free)
128
+ report("Free disk", f"{free_disk:.1f} GB")
129
+ report("Weights need", f"~{WEIGHTS_GB:.2f} GB")
130
+ if free_disk < MIN_DISK_GB:
131
+ problems.append(
132
+ f"{free_disk:.1f} GB free is not enough to cache the weights "
133
+ f"(need ~{MIN_DISK_GB:.0f} GB with headroom)"
134
+ )
135
+
136
+ # --- Verdict -----------------------------------------------------------
137
+ lines.append("")
138
+ for w in warnings:
139
+ lines.append(f" ! {w}")
140
+ for p in problems:
141
+ lines.append(f" x {p}")
142
+ if not warnings and not problems:
143
+ lines.append(" All checks passed.")
144
+
145
+ lines.append("")
146
+ verdict = (
147
+ "BLOCKED — fix the items marked x before loading the model."
148
+ if problems
149
+ else "READY — this machine can load and run Yug."
150
+ )
151
+ lines.append(verdict)
152
+
153
+ if args.quiet:
154
+ print(verdict)
155
+ else:
156
+ print("\n".join(lines))
157
+
158
+ return 1 if problems else 0
159
+
160
+
161
+ if __name__ == "__main__":
162
+ raise SystemExit(main())
yug/base.py ADDED
@@ -0,0 +1,171 @@
1
+ # Copyright 2026 Birla AI Labs
2
+ # SPDX-License-Identifier: Apache-2.0
3
+ """The contract every Yug pipeline obeys, and the type they return.
4
+
5
+ Splitting this out is what lets several checkpoints, sizes, or future variants
6
+ share one user experience: load with ``from_pretrained``, call ``predict``, get
7
+ a :class:`QuantileForecast` back.
8
+ """
9
+
10
+ from __future__ import annotations
11
+
12
+ import abc
13
+ from collections.abc import Sequence
14
+ from typing import Any
15
+
16
+ import numpy as np
17
+
18
+ __all__ = ["ContextLike", "QuantileForecast", "BaseForecastPipeline"]
19
+
20
+ #: Anything accepted as forecast context: one series, several series, or a
21
+ #: covariate stack. Concrete shapes are normalised in the pipeline.
22
+ ContextLike = (
23
+ np.ndarray | Sequence[float] | Sequence[Sequence[float]] | Sequence[np.ndarray]
24
+ )
25
+
26
+
27
+ class QuantileForecast:
28
+ """A probabilistic forecast: quantile levels by horizon, per series.
29
+
30
+ ``values`` is ``(n_series, n_quantiles, horizon)``, ascending in the
31
+ quantile axis. Levels are exactly those the checkpoint was trained to emit
32
+ — the head's width is fixed at training time, so you cannot ask for a level
33
+ that was not trained.
34
+
35
+ Index it to get one series::
36
+
37
+ forecast[0].median # (horizon,)
38
+ forecast[0].quantile(0.9) # (horizon,)
39
+ """
40
+
41
+ def __init__(
42
+ self,
43
+ values: np.ndarray,
44
+ quantile_levels: Sequence[float],
45
+ item_ids: Sequence[Any] | None = None,
46
+ ) -> None:
47
+ values = np.asarray(values, dtype=np.float32)
48
+ if values.ndim == 2:
49
+ values = values[None, ...]
50
+ if values.ndim != 3:
51
+ raise ValueError(
52
+ f"values must be (n_series, n_quantiles, horizon), got {values.shape}"
53
+ )
54
+ if values.shape[1] != len(quantile_levels):
55
+ raise ValueError(
56
+ f"values has {values.shape[1]} quantiles but {len(quantile_levels)} "
57
+ f"levels were given"
58
+ )
59
+
60
+ self.values = values
61
+ self.quantile_levels = [float(q) for q in quantile_levels]
62
+ self.item_ids = list(item_ids) if item_ids is not None else None
63
+
64
+ # ------------------------------------------------------------- geometry
65
+ @property
66
+ def n_series(self) -> int:
67
+ return self.values.shape[0]
68
+
69
+ @property
70
+ def horizon(self) -> int:
71
+ return self.values.shape[2]
72
+
73
+ def __len__(self) -> int:
74
+ return self.n_series
75
+
76
+ def __repr__(self) -> str:
77
+ return (
78
+ f"QuantileForecast(n_series={self.n_series}, horizon={self.horizon}, "
79
+ f"quantile_levels={self.quantile_levels})"
80
+ )
81
+
82
+ def __getitem__(self, i: int) -> QuantileForecast:
83
+ return QuantileForecast(
84
+ self.values[i : i + 1],
85
+ self.quantile_levels,
86
+ None if self.item_ids is None else self.item_ids[i : i + 1],
87
+ )
88
+
89
+ # -------------------------------------------------------------- accessors
90
+ def _level_index(self, q: float) -> int:
91
+ for i, level in enumerate(self.quantile_levels):
92
+ if abs(level - q) < 1e-6:
93
+ return i
94
+ raise ValueError(
95
+ f"quantile {q} was not produced by this checkpoint; available levels "
96
+ f"are {self.quantile_levels}"
97
+ )
98
+
99
+ @property
100
+ def median(self) -> np.ndarray:
101
+ """The 0.5 level, or the nearest level if 0.5 was not trained.
102
+
103
+ ``(horizon,)`` for a single series, ``(n_series, horizon)`` otherwise.
104
+ """
105
+ idx = int(np.argmin(np.abs(np.asarray(self.quantile_levels) - 0.5)))
106
+ out = self.values[:, idx, :]
107
+ return out[0] if self.n_series == 1 else out
108
+
109
+ def quantile(self, q: float) -> np.ndarray:
110
+ """One quantile level across the horizon."""
111
+ out = self.values[:, self._level_index(q), :]
112
+ return out[0] if self.n_series == 1 else out
113
+
114
+ def interval(
115
+ self, lower: float | None = None, upper: float | None = None
116
+ ) -> dict[str, np.ndarray]:
117
+ """A prediction band. Defaults to the widest trained pair."""
118
+ lower = self.quantile_levels[0] if lower is None else lower
119
+ upper = self.quantile_levels[-1] if upper is None else upper
120
+ return {"lower": self.quantile(lower), "upper": self.quantile(upper)}
121
+
122
+ def to_dataframe(self):
123
+ """Long-format frame: one row per (item, step), one column per level.
124
+
125
+ Requires pandas, which is an optional dependency of this package.
126
+ """
127
+ try:
128
+ import pandas as pd
129
+ except ImportError as exc: # pragma: no cover - depends on the install
130
+ raise ImportError(
131
+ "to_dataframe() needs pandas: pip install 'yug[pandas]'"
132
+ ) from exc
133
+
134
+ frames = []
135
+ for i in range(self.n_series):
136
+ item = self.item_ids[i] if self.item_ids is not None else i
137
+ data = {"item_id": item, "step": np.arange(self.horizon)}
138
+ for j, level in enumerate(self.quantile_levels):
139
+ data[str(level)] = self.values[i, j]
140
+ frames.append(pd.DataFrame(data))
141
+ return pd.concat(frames, ignore_index=True)
142
+
143
+
144
+ class BaseForecastPipeline(abc.ABC):
145
+ """Interface shared by every Yug pipeline."""
146
+
147
+ #: Quantile levels this pipeline emits, ascending.
148
+ quantile_levels: list[float]
149
+
150
+ @classmethod
151
+ @abc.abstractmethod
152
+ def from_pretrained(
153
+ cls,
154
+ model_id: str,
155
+ *,
156
+ device_map: str | None = None,
157
+ **kwargs: Any,
158
+ ) -> BaseForecastPipeline:
159
+ """Load a checkpoint by Hub id or local path."""
160
+
161
+ @abc.abstractmethod
162
+ def predict(
163
+ self,
164
+ context: ContextLike,
165
+ prediction_length: int | None = None,
166
+ ) -> QuantileForecast:
167
+ """Forecast ``prediction_length`` steps beyond each context series.
168
+
169
+ Implementations may add keyword-only options on top of these, but must
170
+ accept at least these two positionally.
171
+ """
yug/configs.py ADDED
@@ -0,0 +1,228 @@
1
+ # Copyright 2026 Birla AI Labs
2
+ # SPDX-License-Identifier: Apache-2.0
3
+ """Typed configuration objects for Yug.
4
+
5
+ Three configs, deliberately separated by lifetime:
6
+
7
+ ``YugConfig``
8
+ The *architecture*. Fixed at training time, published as ``config.json``
9
+ next to the weights. Changing any field invalidates a checkpoint.
10
+
11
+ ``InferenceConfig``
12
+ How a batch is fed to the network (device, batch size, workers). Free to
13
+ change per run; has no effect on the numbers the model produces.
14
+
15
+ ``PredictionConfig``
16
+ How the autoregressive rollout is driven (context window, number of sampled
17
+ paths, RNG seed, engine). Affects the forecast, but not the weights.
18
+ """
19
+
20
+ from __future__ import annotations
21
+
22
+ import json
23
+ from dataclasses import asdict, dataclass, field
24
+ from pathlib import Path
25
+ from typing import Any
26
+
27
+ import torch
28
+
29
+ __all__ = ["YugConfig", "InferenceConfig", "PredictionConfig", "default_device"]
30
+
31
+
32
+ def default_device() -> str:
33
+ """``"cuda"`` when an accelerator is present, otherwise ``"cpu"``."""
34
+ return "cuda" if torch.cuda.is_available() else "cpu"
35
+
36
+
37
+ # The exact key set the architecture reads. Anything else in a config.json is
38
+ # metadata (``model_type``, ``architectures``, ``transformers_version``, ...)
39
+ # and is carried alongside rather than passed into the modules.
40
+ _ARCH_KEYS = (
41
+ "patch_len",
42
+ "output_dim",
43
+ "hidden_dim",
44
+ "d_model",
45
+ "freq_num",
46
+ "num_attn_heads",
47
+ "dropout",
48
+ "eps",
49
+ "kernel_size",
50
+ "num_decoder_layers",
51
+ "expansion",
52
+ "quantiles",
53
+ )
54
+
55
+
56
+ @dataclass
57
+ class YugConfig:
58
+ """Architecture hyper-parameters. Serialised as the Hub's ``config.json``.
59
+
60
+ ``quantiles`` is not a runtime choice: the output head emits
61
+ ``output_dim * len(quantiles)`` values per patch, so its width is frozen
62
+ when the checkpoint is trained.
63
+
64
+ ``kernel_size`` is carried for checkpoint-metadata fidelity; the current
65
+ decoder does not read it.
66
+
67
+ ``expansion`` is the MLP hidden-dim multiplier. It may be fractional (the
68
+ V19 checkpoint trains at 1.25); :class:`~yug.model.Yug_MLP` rounds
69
+ ``hidden_size * expansion`` to the nearest int for the actual layer width.
70
+ """
71
+
72
+ patch_len: int = 16
73
+ output_dim: int = 64
74
+ hidden_dim: int = 960
75
+ d_model: int = 960
76
+ freq_num: int = 9
77
+ num_attn_heads: int = 8
78
+ dropout: float = 0.1
79
+ eps: float = 1e-6
80
+ kernel_size: int = 3
81
+ num_decoder_layers: int = 12
82
+ expansion: float = 1.25
83
+ quantiles: list[float] = field(
84
+ default_factory=lambda: [0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9]
85
+ )
86
+
87
+ model_type: str = "yug"
88
+
89
+ def __post_init__(self) -> None:
90
+ self.quantiles = [float(q) for q in self.quantiles]
91
+ if not self.quantiles:
92
+ raise ValueError("quantiles must be a non-empty list")
93
+ if self.quantiles != sorted(self.quantiles):
94
+ raise ValueError(f"quantiles must be ascending, got {self.quantiles}")
95
+ if self.output_dim % self.patch_len:
96
+ raise ValueError(
97
+ f"output_dim ({self.output_dim}) must be a multiple of "
98
+ f"patch_len ({self.patch_len})"
99
+ )
100
+ if self.d_model % self.num_attn_heads:
101
+ raise ValueError(
102
+ f"d_model ({self.d_model}) must be divisible by "
103
+ f"num_attn_heads ({self.num_attn_heads})"
104
+ )
105
+
106
+ # ------------------------------------------------------------- derived
107
+ @property
108
+ def output_patches(self) -> int:
109
+ """Patches emitted per forward pass (``output_dim // patch_len``)."""
110
+ return self.output_dim // self.patch_len
111
+
112
+ @property
113
+ def num_quantiles(self) -> int:
114
+ return len(self.quantiles)
115
+
116
+ @property
117
+ def head_width(self) -> int:
118
+ """Values the output head emits per input patch."""
119
+ return self.output_patches * self.num_quantiles * self.patch_len
120
+
121
+ # --------------------------------------------------------- (de)serialise
122
+ def to_dict(self) -> dict[str, Any]:
123
+ """Full dict, including metadata keys. Written to ``config.json``."""
124
+ return asdict(self)
125
+
126
+ def to_arch_dict(self) -> dict[str, Any]:
127
+ """Only the keys the network modules consume.
128
+
129
+ This is the dict handed to :class:`~yug.model.Yug_Model`;
130
+ keeping it exact is what makes released checkpoints loadable.
131
+ """
132
+ return {k: getattr(self, k) for k in _ARCH_KEYS}
133
+
134
+ @classmethod
135
+ def from_dict(cls, data: dict[str, Any]) -> YugConfig:
136
+ """Build from a dict, ignoring keys this version does not know about."""
137
+ known = set(cls.__dataclass_fields__)
138
+ return cls(**{k: v for k, v in data.items() if k in known})
139
+
140
+ @classmethod
141
+ def from_json_file(cls, path: str | Path) -> YugConfig:
142
+ with open(path, encoding="utf-8") as fh:
143
+ return cls.from_dict(json.load(fh))
144
+
145
+ def to_json_file(self, path: str | Path, indent: int = 2) -> None:
146
+ with open(path, "w", encoding="utf-8") as fh:
147
+ json.dump(self.to_dict(), fh, indent=indent)
148
+ fh.write("\n")
149
+
150
+
151
+ @dataclass
152
+ class InferenceConfig:
153
+ """Batching and placement. Does not change the forecast values.
154
+
155
+ ``pad_token`` is the sentinel the batching path writes into trailing
156
+ positions; those positions are masked out, so its value never reaches the
157
+ loss or the output — it only has to be distinguishable.
158
+ """
159
+
160
+ patch_len: int = 16
161
+ output_len: int = 64
162
+ pad_token: float = 111181.0
163
+ batch_size: int = 32
164
+ num_workers: int = 0
165
+ quantiles: list[float] = field(default_factory=lambda: [0.1, 0.3, 0.5, 0.9])
166
+ device: str = field(default_factory=default_device)
167
+
168
+ @property
169
+ def output_patches(self) -> int:
170
+ return self.output_len // self.patch_len
171
+
172
+ @classmethod
173
+ def from_model_config(cls, config: YugConfig, **overrides: Any) -> InferenceConfig:
174
+ """Derive the batching config from an architecture config."""
175
+ base: dict[str, Any] = {
176
+ "patch_len": config.patch_len,
177
+ "output_len": config.output_dim,
178
+ "quantiles": list(config.quantiles),
179
+ }
180
+ base.update(overrides)
181
+ return cls(**base)
182
+
183
+
184
+ @dataclass
185
+ class PredictionConfig:
186
+ """Rollout behaviour for :meth:`YugPipeline.predict`.
187
+
188
+ ``context_length``
189
+ How many of the most recent points are fed to the model; longer
190
+ histories are cut to their latest ``context_length`` points. Defaults
191
+ to 2048, the window length Yug was trained on and the setting its
192
+ published benchmark scores were produced with. Feeding more history
193
+ is possible but takes the model outside its training regime and, on
194
+ GIFT-Eval, scores worse. ``None`` or ``0`` disables the cut.
195
+
196
+ ``num_samples``
197
+ Independent sampled trajectories. The returned quantile band is the
198
+ empirical quantile across these paths, so raising it narrows Monte
199
+ Carlo noise at linear cost.
200
+
201
+ ``engine``
202
+ ``"cached"`` encodes the context once and reuses its per-layer K/V,
203
+ pushing only newly generated patches through the network (~9x faster).
204
+ ``"exact"`` reproduces the reference implementation's tensor shapes and
205
+ is bitwise identical to it. See :mod:`yug.engine`.
206
+
207
+ ``pad_side``
208
+ Series are grown to a whole number of patches with masked sentinels.
209
+ ``"left"`` keeps the forecast origin on the final real observation and
210
+ is the only side you should normally use; ``"right"`` places the origin
211
+ ``pad_len`` steps past the last real point.
212
+ """
213
+
214
+ context_length: int | None = 2048
215
+ prediction_length: int | None = None
216
+ num_samples: int = 100
217
+ engine: str = "cached"
218
+ pad_side: str = "left"
219
+ seed: int = 0
220
+ quantile_levels: list[float] | None = None
221
+
222
+ def __post_init__(self) -> None:
223
+ if self.engine not in ("cached", "exact"):
224
+ raise ValueError(f"engine must be 'cached' or 'exact', got {self.engine!r}")
225
+ if self.pad_side not in ("left", "right"):
226
+ raise ValueError(f"pad_side must be 'left' or 'right', got {self.pad_side!r}")
227
+ if self.num_samples < 1:
228
+ raise ValueError(f"num_samples must be >= 1, got {self.num_samples}")