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 +5 -0
- yug/__init__.py +57 -0
- yug/__main__.py +162 -0
- yug/base.py +171 -0
- yug/configs.py +228 -0
- yug/engine.py +416 -0
- yug/forecaster.py +485 -0
- yug/gluonts_adapter.py +168 -0
- yug/model.py +997 -0
- yug/pipeline.py +608 -0
- yug/utils.py +239 -0
- yug-0.1.0.dist-info/METADATA +224 -0
- yug-0.1.0.dist-info/RECORD +17 -0
- yug-0.1.0.dist-info/WHEEL +4 -0
- yug-0.1.0.dist-info/entry_points.txt +2 -0
- yug-0.1.0.dist-info/licenses/LICENSE +201 -0
- yug-0.1.0.dist-info/licenses/NOTICE +15 -0
yug/__about__.py
ADDED
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}")
|