views_r2darts2 0.1.1__tar.gz
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.
- views_r2darts2-0.1.1/PKG-INFO +476 -0
- views_r2darts2-0.1.1/README.md +461 -0
- views_r2darts2-0.1.1/pyproject.toml +20 -0
- views_r2darts2-0.1.1/views_r2darts2/__init__.py +21 -0
- views_r2darts2-0.1.1/views_r2darts2/catalogs/README.md +110 -0
- views_r2darts2-0.1.1/views_r2darts2/catalogs/__init__.py +0 -0
- views_r2darts2-0.1.1/views_r2darts2/catalogs/loss_catalog.py +115 -0
- views_r2darts2-0.1.1/views_r2darts2/catalogs/model_catalog.py +353 -0
- views_r2darts2-0.1.1/views_r2darts2/catalogs/optimizer_catalog.py +58 -0
- views_r2darts2-0.1.1/views_r2darts2/catalogs/scheduler_catalog.py +115 -0
- views_r2darts2-0.1.1/views_r2darts2/engines/README.md +101 -0
- views_r2darts2-0.1.1/views_r2darts2/engines/__init__.py +0 -0
- views_r2darts2-0.1.1/views_r2darts2/engines/darts_forecaster.py +908 -0
- views_r2darts2-0.1.1/views_r2darts2/engines/darts_forecasting_model_manager.py +569 -0
- views_r2darts2-0.1.1/views_r2darts2/infrastructure/README.md +117 -0
- views_r2darts2-0.1.1/views_r2darts2/infrastructure/__init__.py +0 -0
- views_r2darts2-0.1.1/views_r2darts2/infrastructure/callbacks.py +1160 -0
- views_r2darts2-0.1.1/views_r2darts2/infrastructure/encoders.py +115 -0
- views_r2darts2-0.1.1/views_r2darts2/infrastructure/exceptions.py +35 -0
- views_r2darts2-0.1.1/views_r2darts2/infrastructure/patches.py +608 -0
- views_r2darts2-0.1.1/views_r2darts2/infrastructure/reproducibility_gate.py +681 -0
- views_r2darts2-0.1.1/views_r2darts2/math/README.md +159 -0
- views_r2darts2-0.1.1/views_r2darts2/math/__init__.py +37 -0
- views_r2darts2-0.1.1/views_r2darts2/math/asymmetric_quantile_loss.py +67 -0
- views_r2darts2-0.1.1/views_r2darts2/math/charbonnier_loss.py +39 -0
- views_r2darts2-0.1.1/views_r2darts2/math/prism_loss.py +384 -0
- views_r2darts2-0.1.1/views_r2darts2/math/sentinel_loss.py +228 -0
- views_r2darts2-0.1.1/views_r2darts2/math/shrinkage_loss.py +60 -0
- views_r2darts2-0.1.1/views_r2darts2/math/spike_focal_loss.py +55 -0
- views_r2darts2-0.1.1/views_r2darts2/math/spotlight_focal_loss.py +235 -0
- views_r2darts2-0.1.1/views_r2darts2/math/spotlight_loss.py +383 -0
- views_r2darts2-0.1.1/views_r2darts2/math/spotlight_loss_asinh.py +409 -0
- views_r2darts2-0.1.1/views_r2darts2/math/spotlight_loss_huber.py +292 -0
- views_r2darts2-0.1.1/views_r2darts2/math/spotlight_loss_logcosh.py +411 -0
- views_r2darts2-0.1.1/views_r2darts2/math/spotlight_loss_power_law.py +278 -0
- views_r2darts2-0.1.1/views_r2darts2/math/time_aware_weighted_huber_loss.py +69 -0
- views_r2darts2-0.1.1/views_r2darts2/math/tweedie_loss.py +114 -0
- views_r2darts2-0.1.1/views_r2darts2/math/warmup_cawr.py +106 -0
- views_r2darts2-0.1.1/views_r2darts2/math/warmup_cosine.py +100 -0
- views_r2darts2-0.1.1/views_r2darts2/math/weighted_huber_loss.py +58 -0
- views_r2darts2-0.1.1/views_r2darts2/math/weighted_penalty_huber_loss.py +120 -0
- views_r2darts2-0.1.1/views_r2darts2/math/zero_inflated_loss.py +78 -0
- views_r2darts2-0.1.1/views_r2darts2/transformers/README.md +172 -0
- views_r2darts2-0.1.1/views_r2darts2/transformers/__init__.py +0 -0
- views_r2darts2-0.1.1/views_r2darts2/transformers/feature_scaler_manager.py +244 -0
- views_r2darts2-0.1.1/views_r2darts2/transformers/scaler_selector.py +263 -0
- views_r2darts2-0.1.1/views_r2darts2/transformers/views_dataset_darts.py +311 -0
|
@@ -0,0 +1,476 @@
|
|
|
1
|
+
Metadata-Version: 2.1
|
|
2
|
+
Name: views_r2darts2
|
|
3
|
+
Version: 0.1.1
|
|
4
|
+
Summary:
|
|
5
|
+
Author: Dylan Pinheiro
|
|
6
|
+
Author-email: dylpin@prio.org
|
|
7
|
+
Requires-Python: >=3.11,<3.15
|
|
8
|
+
Classifier: Programming Language :: Python :: 3
|
|
9
|
+
Classifier: Programming Language :: Python :: 3.11
|
|
10
|
+
Requires-Dist: darts (==0.40.0)
|
|
11
|
+
Requires-Dist: scikit-learn (>=1.6.0,<1.8.0)
|
|
12
|
+
Requires-Dist: views-pipeline-core (>=2.0.0,<3.0.0)
|
|
13
|
+
Description-Content-Type: text/markdown
|
|
14
|
+
|
|
15
|
+
<p align="center">
|
|
16
|
+
<img src="https://github.com/user-attachments/assets/4cd8129b-9ad6-4fa3-a4ca-8288b0ab610f" alt="r2darts2 Banner" width="85%">
|
|
17
|
+
</p>
|
|
18
|
+
<p align="center">
|
|
19
|
+
<img src="https://img.shields.io/badge/python-3.11%2B-blue.svg" alt="Python Version" />
|
|
20
|
+
|
|
21
|
+
<img src="https://img.shields.io/badge/darts-0.40.0-green.svg" alt="Darts Version" />
|
|
22
|
+
|
|
23
|
+
<img src="https://img.shields.io/badge/pytorch-2.x-orange.svg" alt="PyTorch" />
|
|
24
|
+
|
|
25
|
+
<img src="https://img.shields.io/badge/license-MIT-lightgrey.svg" alt="License" />
|
|
26
|
+
</p>
|
|
27
|
+
|
|
28
|
+
<p align="center">
|
|
29
|
+
A time series forecasting package for conflict prediction within the VIEWS (Violence and Impacts Early-Warning System) ecosystem. Built on Darts and PyTorch Lightning, it provides deep learning models, domain-specific loss functions, and reproducibility infrastructure tuned for zero-inflated, heavy-tailed conflict fatality data.
|
|
30
|
+
</p>
|
|
31
|
+
|
|
32
|
+
---
|
|
33
|
+
|
|
34
|
+
## Key Features
|
|
35
|
+
|
|
36
|
+
- ๐ **Production-Ready Integration**: Seamlessly integrates with the VIEWS pipeline ecosystem via the Genomic Firewall and DNA manifest validation
|
|
37
|
+
- โก **Zero-Inflated Data Handling**: Specialized loss functions (SpotlightLoss family) and scalers (AsinhTransform chains) purpose-built for conflict fatality distributions
|
|
38
|
+
- ๐ง **10 Model Architectures**: TFT, N-BEATS, N-HiTS, TiDE, TCN, BlockRNN, Transformer, NLinear, DLinear, TSMixer
|
|
39
|
+
- ๐ **Static Covariate Fingerprints**: Per-entity conflict statistics (ยต, ฯ, max, trend, sparsity)
|
|
40
|
+
- ๐ **Chained Scalers**: Arrow-syntax pipelines (`AsinhTransform->MaxAbsScaler`) for multi-stage feature normalization
|
|
41
|
+
- ๐ก๏ธ **Fortress Architecture**: ADR-governed reproducibility, NaN detection, gradient health monitoring, and training stability callbacks throughout
|
|
42
|
+
- ๐งฎ **GPU Acceleration**: Optimized for single- and multi-GPU training via PyTorch Lightning
|
|
43
|
+
|
|
44
|
+
---
|
|
45
|
+
|
|
46
|
+
## ๐ฆ Installation
|
|
47
|
+
|
|
48
|
+
```bash
|
|
49
|
+
git clone https://github.com/views-platform/views-r2darts2.git
|
|
50
|
+
cd views-r2darts2
|
|
51
|
+
pip install -e .
|
|
52
|
+
```
|
|
53
|
+
|
|
54
|
+
Requires `darts==0.40.0` and `views-pipeline-core>=2.0.0`. For GPU support, install the appropriate PyTorch version for your CUDA setup first. See the [PyTorch installation guide](https://pytorch.org/get-started/locally/).
|
|
55
|
+
|
|
56
|
+
---
|
|
57
|
+
|
|
58
|
+
## ๐ง Supported Models
|
|
59
|
+
|
|
60
|
+
| Model | Static Covariates | Description | Ideal For |
|
|
61
|
+
|-------|:-----------------:|-------------|-----------|
|
|
62
|
+
| **TFT** (Temporal Fusion Transformer) | โ
VSN+GRN gating | Hybrid LSTM + multi-head attention. Variable Selection Networks gate static covariate influence per time step. | Interpretable multivariate forecasting; best when entity identity strongly conditions the forecast. |
|
|
63
|
+
| **TSMixer** | โ
Concatenation | Alternating time-mixing and feature-mixing MLP blocks. Static covariates are concatenated at each block (no gating). | Large-scale multivariate; fast training. Requires `AsinhTransform->MaxAbsScaler` on static cov stats due to blunt concat. |
|
|
64
|
+
| **TiDE** | โ
Concatenation | MLP encoder-decoder for long-horizon forecasting. Efficient and scalable. | Resource-constrained environments, long input horizons. |
|
|
65
|
+
| **BlockRNN** | โ
Concatenation | Stacked RNN/LSTM/GRU with optional static covariate injection. | Sequential dependency modeling; autoregressive forecasting. |
|
|
66
|
+
| **NLinear** | โ
Concatenation | Lightweight linear model with optional trend/seasonality decomposition. | Baseline modeling; rapid prototyping. |
|
|
67
|
+
| **DLinear** | โ
Concatenation | Decomposition-Linear โ separate linear layers for trend and seasonal components. | Trend/seasonality separation; fast inference. |
|
|
68
|
+
| **N-HiTS** | โ Not consumed | Hierarchical interpolation with per-stack poolingโFCโtheta pipelines. Flag accepted but **not passed to the Darts constructor** โ static covariate fingerprints are computed and attached but silently ignored by the model. | Multi-scale temporal pattern extraction; long-term forecasting. |
|
|
69
|
+
| **Transformer** | โ Not consumed | Self-attention encoder with positional encoding. Flag accepted but **silently ignored** โ no `use_static_covariates` is passed to the Darts constructor. | Long-range temporal dependency modeling. |
|
|
70
|
+
| **N-BEATS** | โ Not supported | Fully-connected stacks with basis expansions for trend/seasonality. Static covariates are architecturally incompatible. Do not configure `use_static_covariates=True`. | Interpretable decomposition; univariate/multivariate forecasting. |
|
|
71
|
+
| **TCN** (Temporal Convolutional Network) | โ Not consumed | Dilated causal convolutions with residual connections. | High-frequency series; long-range dependencies. |
|
|
72
|
+
|
|
73
|
+
> **Important**: Models marked โ will silently drop static covariate fingerprints even if `use_static_covariates=True` is set in the config. This is a Darts library constraint, not a bug in this package. For N-HiTS and Transformer, the fingerprints are computed (respecting `stat_time_range` for leakage prevention) but never seen by the model weights.
|
|
74
|
+
|
|
75
|
+
---
|
|
76
|
+
|
|
77
|
+
## ๐ Scalers
|
|
78
|
+
|
|
79
|
+
Proper data scaling is critical for neural network training. For conflict data โ zero-inflated, heavily right-skewed, spanning four orders of magnitude โ the choice of transform directly determines whether models converge.
|
|
80
|
+
|
|
81
|
+
### Available Scalers
|
|
82
|
+
|
|
83
|
+
| Scaler | Formula | Best For | Notes |
|
|
84
|
+
|--------|---------|----------|-------|
|
|
85
|
+
| **AsinhTransform** | $y = \text{asinh}(x)$ | Zero-inflated counts, data with negatives | โญ **Recommended for all fatality data**. Handles zeros and extreme outliers. Near-linear below 1. |
|
|
86
|
+
| **MaxAbsScaler** | $y = x / \max(\|x\|)$ | Cross-entity normalization after element-wise transform | Maps to $[-1,1]$. Preserves ordinal rank. Required after AsinhTransform for static covariate injection in concatenation-based models. |
|
|
87
|
+
| **StandardScaler** | $y = (x - \mu)/\sigma$ | Roughly normal data | Zero-mean, unit-variance. Poor choice for zero-inflated distributions. |
|
|
88
|
+
| **MinMaxScaler** | $y = (x - x_{\min})/(x_{\max} - x_{\min})$ | Bounded data (0โ1, 0โ100) | Best for V-Dem indices and WDI percentages. |
|
|
89
|
+
| **SqrtTransform** | $y = \sqrt{x}$ | Moderate skew, count data | Gentler compression than asinh; undefined for negatives. |
|
|
90
|
+
| **LogTransform** | $y = \log(1 + x)$ | Strictly positive skewed data | Undefined for negatives. Prefer AsinhTransform for conflict data. |
|
|
91
|
+
| **RobustScaler** | Median + IQR | Outlier-heavy data | Resistant to extreme values. |
|
|
92
|
+
| **QuantileNormal** | Maps to $\mathcal{N}(0,1)$ | Any distribution | Forces Gaussian marginal. |
|
|
93
|
+
| **QuantileUniform** | Maps to $U(0,1)$ | Any distribution | Forces uniform marginal. |
|
|
94
|
+
| **YeoJohnsonTransform** | Power transform | Mixed positive/negative | Makes data more Gaussian-like. |
|
|
95
|
+
|
|
96
|
+
### AsinhTransform vs LogTransform
|
|
97
|
+
|
|
98
|
+
```
|
|
99
|
+
Log(1+x) AsinhTransform
|
|
100
|
+
x = -50 โ undefined โ asinh(-50) = -4.61
|
|
101
|
+
x = 0 0.00 0.00
|
|
102
|
+
x = 1 0.69 0.88 (non-zero threshold in asinh space: 0.88 โ 1 death)
|
|
103
|
+
x = 100 4.62 5.30
|
|
104
|
+
x = 10000 9.21 9.90
|
|
105
|
+
```
|
|
106
|
+
|
|
107
|
+
The `non_zero_threshold` in SpotlightLoss is set to `0.88` because `asinh(1) โ 0.88` โ this exactly corresponds to the boundary of "at least 1 battle death" in raw space.
|
|
108
|
+
|
|
109
|
+
### ๐ Chained Scalers
|
|
110
|
+
|
|
111
|
+
Use the `->` operator to compose transforms sequentially. This is the **production standard** for all conflict and covariate features:
|
|
112
|
+
|
|
113
|
+
```python
|
|
114
|
+
# Target: suppress zeros + bound output
|
|
115
|
+
"target_scaler": "AsinhTransform"
|
|
116
|
+
|
|
117
|
+
# Features: element-wise compression then cross-entity normalization
|
|
118
|
+
"feature_scaler_map": {
|
|
119
|
+
"AsinhTransform->MaxAbsScaler": [
|
|
120
|
+
"lr_splag_1_ged_sb", "lr_ged_ns", "lr_ged_os",
|
|
121
|
+
"lr_acled_sb", "lr_wdi_ny_gdp_mktp_kd",
|
|
122
|
+
# ... all conflict and macro features
|
|
123
|
+
],
|
|
124
|
+
}
|
|
125
|
+
|
|
126
|
+
# Static covariate statistics (for models that consume them)
|
|
127
|
+
"static_covariate_stats": {"transform": "AsinhTransform->MaxAbsScaler"}
|
|
128
|
+
```
|
|
129
|
+
|
|
130
|
+
**Why MaxAbsScaler after AsinhTransform for features?**
|
|
131
|
+
AsinhTransform is element-wise: it compresses Syria's `ged_sb โ 5000` and Chad's `ged_sb โ 3` independently. After asinh, Syria is at ~8.5 and Chad at ~1.8 โ a 4.7ร gap persists. MaxAbsScaler maps the entire feature column to `[-1, 1]` across all 180+ countries, collapsing cross-entity scale while preserving ordinal rank. Without this, concatenation-based models (TSMixer, TiDE) inject raw magnitude bias at every block.
|
|
132
|
+
|
|
133
|
+
**Forward / inverse chain direction:**
|
|
134
|
+
```
|
|
135
|
+
Forward: X โ Scalerโ.fit_transform(X) โ Scalerโ.fit_transform(X') โ X_scaled
|
|
136
|
+
Inverse: X_scaled โ Scalerโ.inverse_transform โ Scalerโ.inverse_transform โ X_original
|
|
137
|
+
```
|
|
138
|
+
|
|
139
|
+
### Static Covariate Fingerprints
|
|
140
|
+
|
|
141
|
+
For models that consume static covariates (TFT, TSMixer, TiDE, BlockRNN, NLinear, DLinear), five per-entity statistics are computed from the **training partition only** (via `stat_time_range`) and injected as `TimeSeries.static_covariates` metadata:
|
|
142
|
+
|
|
143
|
+
| Statistic | Meaning | Transform Recommendation |
|
|
144
|
+
|-----------|---------|--------------------------|
|
|
145
|
+
| `target_mu` | Mean conflict level | `AsinhTransform->MaxAbsScaler` |
|
|
146
|
+
| `target_sigma` | Volatility / spread | `AsinhTransform->MaxAbsScaler` |
|
|
147
|
+
| `target_max` | Peak value (spike extremity) | `AsinhTransform->MaxAbsScaler` |
|
|
148
|
+
| `target_trend` | OLS slope over training window | `AsinhTransform->MaxAbsScaler` |
|
|
149
|
+
| `target_sparsity` | Fraction of zero months | No transform (already in [0,1]) |
|
|
150
|
+
|
|
151
|
+
Always pass `stat_time_range` to `as_darts_timeseries()` to prevent test-period leakage:
|
|
152
|
+
```python
|
|
153
|
+
ts_list = dataset.as_darts_timeseries(
|
|
154
|
+
stat_time_range=(training_start_month_id, training_end_month_id),
|
|
155
|
+
static_cov_transform="AsinhTransform->MaxAbsScaler",
|
|
156
|
+
)
|
|
157
|
+
```
|
|
158
|
+
|
|
159
|
+
### Recommended Scaler by Data Source
|
|
160
|
+
|
|
161
|
+
| Data Source | Feature Type | Recommended Chain | Rationale |
|
|
162
|
+
|-------------|--------------|-------------------|-----------|
|
|
163
|
+
| **UCDP / ACLED** | Fatality counts | `AsinhTransform->MaxAbsScaler` | Zero-inflated; cross-entity normalization required |
|
|
164
|
+
| **WDI** | GDP, population, aid flows | `AsinhTransform->MaxAbsScaler` | Spans many orders of magnitude; can be negative (net migration) |
|
|
165
|
+
| **WDI** | Percentages (`_zs` suffix) | `AsinhTransform->MaxAbsScaler` | Handles near-zero values; consistent with other features |
|
|
166
|
+
| **V-Dem** | Democracy indices (0โ1 bounded) | `AsinhTransform->MaxAbsScaler` or `MinMaxScaler` | Already bounded; either works |
|
|
167
|
+
| **Static cov stats** | ยต, ฯ, max, trend | `AsinhTransform->MaxAbsScaler` | Required for concatenation-based models |
|
|
168
|
+
| **Static cov stats** | Sparsity | None (raw) | Already in [0,1] |
|
|
169
|
+
|
|
170
|
+
---
|
|
171
|
+
|
|
172
|
+
## โก Loss Functions
|
|
173
|
+
|
|
174
|
+
All loss functions target **zero-inflated conflict data**: ~90% zeros, ~10% events spanning four orders of magnitude. The loss function family has evolved substantially โ the table below shows the current production-recommended functions and the full catalog.
|
|
175
|
+
|
|
176
|
+
### Loss Function Catalog
|
|
177
|
+
|
|
178
|
+
| Loss Function | Status | Base Cell Loss | Key Mechanism | Use When |
|
|
179
|
+
|---------------|--------|---------------|---------------|----------|
|
|
180
|
+
| **SpotlightLossLogcosh** | โญ **Production** | log_cosh | DC/AC decomp + compound weights + KL-DRO + level anchor + spectral | Default for all models in production |
|
|
181
|
+
| **SpotlightLoss** | โญ **Production** | Barron(ฮฑ=1.5) | Same as above with more robust base cell loss | When log_cosh gradient is too aggressive on large errors |
|
|
182
|
+
| **PrismLoss** | Research | MSE (= MSLE in log space) | KL-DRO + compound weights, no DC/AC decomp, no level anchor | MSLE-aligned optimization without RevIN |
|
|
183
|
+
| **SpotlightFocalLoss** | Research | log_cosh | Focal weighting by difficulty `(1โexp(โ\|e\|))^ฮณ`, no DRO | Models without RevIN; exploration |
|
|
184
|
+
| **SentinelLoss** | Research | Generalised Charbonnier | Power-law magnitude weights + SiLU symmetry + temporal gradient | Alternative robust base when Barron ฮฑ needs tuning |
|
|
185
|
+
| **WeightedPenaltyHuberLoss** | Legacy | Huber | FP/FN multiplicative penalties | Simple baselines; not recommended for production |
|
|
186
|
+
| **WeightedHuberLoss** | Legacy | Huber | Non-zero reweighting | Simple baselines |
|
|
187
|
+
| **TimeAwareWeightedHuberLoss** | Legacy | Huber | Temporal decay + event weights | Time-sensitive ablations |
|
|
188
|
+
| **TweedieLoss** | Legacy | Tweedie (pโ1.5) | Compound Poisson-Gamma | Count data without asinh transform |
|
|
189
|
+
| **AsymmetricQuantileLoss** | Legacy | Quantile | Asymmetric ฯ-penalty | When underestimation cost >> overestimation |
|
|
190
|
+
| **ZeroInflatedLoss** | Legacy | Huber (two-part) | Explicit binary + count split | Explicit zero-inflation modeling |
|
|
191
|
+
| **SpikeFocalLoss** | Legacy | log_cosh | Focal on absolute magnitude | Predates KL-DRO; superseded by Spotlight family |
|
|
192
|
+
| **ShrinkageLoss** | Legacy | Shrinkage | Suppresses easy samples via sigmoid gate | Exploratory; not validated for conflict |
|
|
193
|
+
|
|
194
|
+
### SpotlightLossLogcosh โ Architecture Deep Dive
|
|
195
|
+
|
|
196
|
+
The production loss for all current VIEWS models. Operates entirely in **asinh space**; the target scaler must be `AsinhTransform`.
|
|
197
|
+
|
|
198
|
+
**Five orthogonal components:**
|
|
199
|
+
|
|
200
|
+
**1. DC/AC decomposition** โ prevents RevIN from amplifying bias:
|
|
201
|
+
```
|
|
202
|
+
e_shape = e โ mean(e) per series
|
|
203
|
+
```
|
|
204
|
+
The shape gradient sums to zero per series by construction (`J = I โ 11แต/T`). A small bias `b` in normalized space becomes `bยทฯ` after RevIN denormalization, and `sinh(bยทฯ) > sinh(E[bยทฯ])` via Jensen's inequality โ exponential overprediction in raw death counts. The DC/AC split structurally blocks this. The level anchor (component 4) is the *only* mechanism that can shift per-series means.
|
|
205
|
+
|
|
206
|
+
**2. Adaptive compound weighting** โ parameter-free event focus:
|
|
207
|
+
```
|
|
208
|
+
difficulty = 1 โ exp(โ|e_shape|) โ [0, 1)
|
|
209
|
+
importance = 1 โ exp(โmax(|y|, |ลท_sg|)) โ [0, 1)
|
|
210
|
+
w_compound = 1 + difficulty ร importance โ [1, 2)
|
|
211
|
+
```
|
|
212
|
+
Both signals must be active simultaneously. Perfect predictions get `difficultyโ0 โ wโ1` regardless of magnitude. Replaces the `alpha` hyperparameter from earlier versions.
|
|
213
|
+
|
|
214
|
+
**3. KL-DRO tail aggregation** โ proportional outlier detection:
|
|
215
|
+
```
|
|
216
|
+
log_l = log(l + ฮต)
|
|
217
|
+
z = (log_l โ mean(log_l)) / std(log_l)
|
|
218
|
+
dro_w = log1p(clamp(1+z, min=0))
|
|
219
|
+
dro_w = dro_w / mean(dro_w)
|
|
220
|
+
ฮฑ_soft = log_std / (log_std + 1.0) # soft activation (uniform early in training)
|
|
221
|
+
```
|
|
222
|
+
Unlike ฯยฒ-DRO (which detects *absolute* loss outliers and causes Syria to dominate), KL-DRO detects *proportional* outliers: a village miss at 10ร the median receives the same weight as a Syria miss at 10ร the median. Aligned with the proportional error sensitivity of asinh-space MSE.
|
|
223
|
+
|
|
224
|
+
**4. Level anchor** โ T-scaled log_cosh on per-series mean error:
|
|
225
|
+
```
|
|
226
|
+
L_level = T ยท mean_per_series[ log_cosh(mean(ลท) โ mean(y)) ]
|
|
227
|
+
```
|
|
228
|
+
The *only* mechanism that can shift series-level means. T-scaling compensates for the `1/T` chain-rule factor from mean reduction.
|
|
229
|
+
|
|
230
|
+
**5. Spectral regularization** (optional, `ฮด > 0`):
|
|
231
|
+
Multi-resolution STFT magnitude comparison with the DC bin masked. Enforces temporal structure without caring about phase. Proportional contribution tuned via `delta` (current production values: 0.015โ0.12 depending on model).
|
|
232
|
+
|
|
233
|
+
**Configuration:**
|
|
234
|
+
```python
|
|
235
|
+
"loss_function": "SpotlightLossLogcosh",
|
|
236
|
+
"delta": 0.02, # Spectral weight; 0 disables spectral term
|
|
237
|
+
"non_zero_threshold": 0.88, # asinh(1) โ 0.88 (= 1 battle death in raw space)
|
|
238
|
+
```
|
|
239
|
+
|
|
240
|
+
### Loss Function Evolution
|
|
241
|
+
|
|
242
|
+
The loss function development followed a clear progression as each failure mode was identified and addressed:
|
|
243
|
+
|
|
244
|
+
| Generation | Loss | Problem It Solved | Limitation Discovered |
|
|
245
|
+
|---|---|---|---|
|
|
246
|
+
| Gen 1 | `WeightedPenaltyHuberLoss` | Basic non-zero reweighting | Huber is symmetric; no handling of RevIN bias; manually tuned FP/FN parameters |
|
|
247
|
+
| Gen 2 | `TweedieLoss`, `SpikeFocalLoss` | Compound Poisson structure; focal weighting | No DRO; focal exponent ฮณ is a fragile hyperparameter |
|
|
248
|
+
| Gen 3 | `PrismLoss` | KL-DRO replaces ฯยฒ-DRO; compound weighting is parameter-free | No DC/AC decomposition; RevIN bias accumulates; no level anchor |
|
|
249
|
+
| Gen 4 | `SpotlightFocalLoss` | Focal mechanism adapted for regression, no class-specific logic | No DRO; still requires ฮณ tuning |
|
|
250
|
+
| Gen 5 | `SpotlightLossLogcosh` | DC/AC decomp + compound + KL-DRO + level anchor + spectral; fully parameter-free weighting | log_cosh gradient can clip large errors aggressively |
|
|
251
|
+
| Gen 5b | `SpotlightLoss` | Barron(ฮฑ=1.5) base cell loss โ heavier tail than log_cosh | Same architecture as v36 Logcosh |
|
|
252
|
+
|
|
253
|
+
---
|
|
254
|
+
|
|
255
|
+
## ๐๏ธ Multi-Stack Model Configuration
|
|
256
|
+
|
|
257
|
+
N-HiTS and TSMixer are multi-stack architectures where residuals are passed between stacks. **Layer width ordering is critical and non-obvious.**
|
|
258
|
+
|
|
259
|
+
### The Layer Widths Trap
|
|
260
|
+
|
|
261
|
+
Both N-HiTS and TSMixer use a **residual stacking pipeline**:
|
|
262
|
+
|
|
263
|
+
```
|
|
264
|
+
Stack 0 (coarse) โ absorbs easy patterns (trend, long cycles)
|
|
265
|
+
โ residual
|
|
266
|
+
Stack 1 (mid) โ absorbs medium-frequency patterns
|
|
267
|
+
โ residual
|
|
268
|
+
Stack 2 (fine) โ must absorb ALL remaining residuals, including spike patterns
|
|
269
|
+
```
|
|
270
|
+
|
|
271
|
+
**The fine stack always has the hardest job.** Under SpotlightLoss (DRO weighting), high-conflict country residuals (e.g., Sudan) dominate gradients. If the fine stack has minimal capacity, it cannot model these spikes, producing erratic theta coefficients that cause:
|
|
272
|
+
|
|
273
|
+
- **Explosion on high-conflict countries** (Sudan, Syria): fine stack theta blows up
|
|
274
|
+
- **Flatline on peaceful countries**: coarse stack weights drift toward dominant loss signal, collapsing peaceful predictions to near-zero
|
|
275
|
+
|
|
276
|
+
**Correct configuration:**
|
|
277
|
+
```python
|
|
278
|
+
# WRONG โ coarse gets most capacity, fine gets least
|
|
279
|
+
"layer_widths": [256, 128, 64] # โ Sudan explosion + flatline
|
|
280
|
+
|
|
281
|
+
# CORRECT โ fine stack gets most capacity for residual absorption
|
|
282
|
+
"layer_widths": [64, 128, 256] # โ stable, fine stack can handle spikes
|
|
283
|
+
```
|
|
284
|
+
|
|
285
|
+
### N-HiTS Pooling / Frequency Alignment
|
|
286
|
+
|
|
287
|
+
N-HiTS pools the input sequence before each stack's FC block. The `pooling_kernel_sizes` and `n_freq_downsample` must be **aligned**:
|
|
288
|
+
|
|
289
|
+
```
|
|
290
|
+
pool_k = 4 โ input compressed to ceil(36/4) = 9 time steps โ FC input dim = 9
|
|
291
|
+
n_freq = 4 โ theta output has 4 frequency coefficients, interpolated to 36 steps
|
|
292
|
+
```
|
|
293
|
+
|
|
294
|
+
If `pool_k=4` but `n_freq=3`, there are 9 FC inputs but only 3 theta points โ implicit upsampling by 9โ36 via 3 basis functions creates a 3:1 gap. This forces the interpolation to guess intermediate values, introducing artificial smoothing that conflicts with spike reconstruction.
|
|
295
|
+
|
|
296
|
+
**Correct alignment:**
|
|
297
|
+
```python
|
|
298
|
+
"pooling_kernel_sizes": [[4, 2, 1]], # coarse: 9 steps, mid: 18 steps, fine: 36 steps
|
|
299
|
+
"n_freq_downsample": [[4, 2, 1]], # 4 theta / 2 theta / 1 theta, interpolated to 36
|
|
300
|
+
```
|
|
301
|
+
|
|
302
|
+
The fine stack (`pool_k=1, n_freq=1`) sees all 36 time steps and produces 36 theta coefficients โ effectively identity interpolation. This is correct: the fine stack should not impose any temporal compression on spike signals.
|
|
303
|
+
|
|
304
|
+
Also use `max_pool_1d=True` in the coarse stack to preserve spike maxima during pooling (average pooling dilutes spike information that should be routed to the coarse trend stack, not the fine detail stack).
|
|
305
|
+
|
|
306
|
+
---
|
|
307
|
+
|
|
308
|
+
## โก Loss Functions
|
|
309
|
+
|
|
310
|
+
*(See complete catalog above.)*
|
|
311
|
+
|
|
312
|
+
---
|
|
313
|
+
|
|
314
|
+
## ๐ก๏ธ Fortress Architecture & Governance
|
|
315
|
+
|
|
316
|
+
This repository adheres to the **Fortress Architecture**: strict engineering and mathematical standards designed to guarantee scientific integrity and reproducibility in conflict forecasting.
|
|
317
|
+
|
|
318
|
+
The repository is governed by:
|
|
319
|
+
- **[Architectural Decision Records (ADRs)](docs/ADRs/README.md)**: Sequential, authoritative records of every major design choice.
|
|
320
|
+
- **[Class Intent Contracts (CICs)](docs/CICs/README.md)**: Explicit declarations of purpose and responsibility for every critical class.
|
|
321
|
+
- **[Reproducibility Manifest](docs/standards/REPRODUCIBILITY_MANIFEST.md)**: The mandatory DNA genome that every experiment must declare before execution.
|
|
322
|
+
|
|
323
|
+
### Training Stability Callbacks
|
|
324
|
+
|
|
325
|
+
Every training run is monitored by mandatory Fortress callbacks configured in `ModelCatalog`:
|
|
326
|
+
|
|
327
|
+
| Callback | Purpose |
|
|
328
|
+
|----------|---------|
|
|
329
|
+
| `NaNDetectionCallback` | Halts training immediately on NaN in loss or weights |
|
|
330
|
+
| `GradientHealthCallback` | Monitors gradient norm; warns on explosion/vanishing |
|
|
331
|
+
| `WeightNormCallback` | Tracks parameter norm evolution across epochs |
|
|
332
|
+
| `LossStabilityCallback` | Detects loss spikes and plateau regimes |
|
|
333
|
+
| `RevINMonitorCallback` | Monitors RevIN affine parameters for drift |
|
|
334
|
+
| `PredictionSanityCallback` | Validates output shape and value range each epoch |
|
|
335
|
+
| `YHatBarCallback` | Tracks per-series mean predictions (ลท bar) against targets |
|
|
336
|
+
| `EpochTimingCallback` | Logs wall-clock time per epoch for performance tracking |
|
|
337
|
+
|
|
338
|
+
---
|
|
339
|
+
|
|
340
|
+
## ๐ง API Reference
|
|
341
|
+
|
|
342
|
+
Every core class follows the **1-Class-1-File** standard.
|
|
343
|
+
|
|
344
|
+
### ScalerSelector
|
|
345
|
+
```python
|
|
346
|
+
from views_r2darts2.transformers.scaler_selector import ScalerSelector
|
|
347
|
+
|
|
348
|
+
scaler = ScalerSelector.get_scaler("AsinhTransform")
|
|
349
|
+
pipeline = ScalerSelector.get_chained_scaler("AsinhTransform->MaxAbsScaler")
|
|
350
|
+
```
|
|
351
|
+
|
|
352
|
+
### FeatureScalerManager
|
|
353
|
+
```python
|
|
354
|
+
from views_r2darts2.transformers.feature_scaler_manager import FeatureScalerManager
|
|
355
|
+
|
|
356
|
+
manager = FeatureScalerManager(
|
|
357
|
+
feature_scaler_map={"AsinhTransform->MaxAbsScaler": ["lr_ged_sb", "lr_ged_ns"]},
|
|
358
|
+
default_scaler=None,
|
|
359
|
+
)
|
|
360
|
+
```
|
|
361
|
+
|
|
362
|
+
### Triple Catalogs (Genomic Firewall)
|
|
363
|
+
```python
|
|
364
|
+
from views_r2darts2.catalogs.model_catalog import ModelCatalog
|
|
365
|
+
from views_r2darts2.catalogs.loss_catalog import LossCatalog
|
|
366
|
+
from views_r2darts2.catalogs.optimizer_catalog import OptimizerCatalog
|
|
367
|
+
from views_r2darts2.catalogs.scheduler_catalog import SchedulerCatalog
|
|
368
|
+
|
|
369
|
+
# Catalogs validate the DNA manifest on initialization
|
|
370
|
+
loss_fn = LossCatalog(config).get_loss()
|
|
371
|
+
model = ModelCatalog(config).get_model("NHiTSModel")
|
|
372
|
+
```
|
|
373
|
+
|
|
374
|
+
### _ViewsDatasetDarts
|
|
375
|
+
```python
|
|
376
|
+
from views_r2darts2.transformers.views_dataset_darts import _ViewsDatasetDarts
|
|
377
|
+
|
|
378
|
+
dataset = _ViewsDatasetDarts.from_views_path(path_raw, run_type, config)
|
|
379
|
+
|
|
380
|
+
# Always pass stat_time_range to prevent leakage into static covariate stats
|
|
381
|
+
ts_list = dataset.as_darts_timeseries(
|
|
382
|
+
stat_time_range=(training_start_id, training_end_id),
|
|
383
|
+
static_cov_transform="AsinhTransform->MaxAbsScaler",
|
|
384
|
+
)
|
|
385
|
+
```
|
|
386
|
+
|
|
387
|
+
### ReproducibilityGate
|
|
388
|
+
```python
|
|
389
|
+
from views_r2darts2.infrastructure.reproducibility_gate import ReproducibilityGate
|
|
390
|
+
|
|
391
|
+
ReproducibilityGate.Config.audit_manifest(config)
|
|
392
|
+
ReproducibilityGate.Data.audit_dataframe_schema(df, expected_targets, expected_features)
|
|
393
|
+
ReproducibilityGate.Temporal.audit_continuity(partition)
|
|
394
|
+
```
|
|
395
|
+
|
|
396
|
+
---
|
|
397
|
+
|
|
398
|
+
## ๐ Production Configuration Template
|
|
399
|
+
|
|
400
|
+
Minimal validated configuration for a new model:
|
|
401
|
+
|
|
402
|
+
```python
|
|
403
|
+
def get_hp_config():
|
|
404
|
+
return {
|
|
405
|
+
# Forecast horizon
|
|
406
|
+
"steps": [*range(1, 37)],
|
|
407
|
+
"input_chunk_length": 36,
|
|
408
|
+
"output_chunk_length": 36,
|
|
409
|
+
"output_chunk_shift": 0,
|
|
410
|
+
|
|
411
|
+
# Scaling โ production standard
|
|
412
|
+
"target_scaler": "AsinhTransform",
|
|
413
|
+
"feature_scaler": None,
|
|
414
|
+
"feature_scaler_map": {
|
|
415
|
+
"AsinhTransform->MaxAbsScaler": [
|
|
416
|
+
# All conflict counts, macro indicators, and lagged features
|
|
417
|
+
],
|
|
418
|
+
},
|
|
419
|
+
"static_covariate_stats": {"transform": "AsinhTransform->MaxAbsScaler"},
|
|
420
|
+
|
|
421
|
+
# Loss โ production standard
|
|
422
|
+
"loss_function": "SpotlightLossLogcosh",
|
|
423
|
+
"delta": 0.02, # Spectral weight; tune via W&B sweep
|
|
424
|
+
"non_zero_threshold": 0.88, # asinh(1): boundary of 1 battle death
|
|
425
|
+
|
|
426
|
+
# Optimizer
|
|
427
|
+
"optimizer_cls": "AdamW",
|
|
428
|
+
"lr": 0.0005,
|
|
429
|
+
"weight_decay": 0.0002,
|
|
430
|
+
"gradient_clip_val": 3,
|
|
431
|
+
"optimizer_kwargs": {"lr": 0.0005, "weight_decay": 0.0002},
|
|
432
|
+
|
|
433
|
+
# Scheduler
|
|
434
|
+
"lr_scheduler_cls": "ReduceLROnPlateau",
|
|
435
|
+
"lr_scheduler_factor": 0.5,
|
|
436
|
+
"lr_scheduler_patience": 12,
|
|
437
|
+
"lr_scheduler_min_lr": 1e-6,
|
|
438
|
+
"lr_scheduler_kwargs": {
|
|
439
|
+
"mode": "min", "factor": 0.5, "patience": 12,
|
|
440
|
+
"min_lr": 1e-6, "cooldown": 3,
|
|
441
|
+
"threshold": 0.01, "threshold_mode": "rel",
|
|
442
|
+
},
|
|
443
|
+
|
|
444
|
+
# Training
|
|
445
|
+
"batch_size": 128,
|
|
446
|
+
"n_epochs": 300,
|
|
447
|
+
"early_stopping_patience": 35,
|
|
448
|
+
"early_stopping_min_delta": 0.001,
|
|
449
|
+
"force_reset": True,
|
|
450
|
+
|
|
451
|
+
# Normalization
|
|
452
|
+
"use_reversible_instance_norm": True,
|
|
453
|
+
"use_cyclic_encoders": True,
|
|
454
|
+
"use_static_covariates": True, # Set False for N-BEATS, N-HiTS, Transformer
|
|
455
|
+
|
|
456
|
+
# Reproducibility
|
|
457
|
+
"random_state": 67,
|
|
458
|
+
"time_steps": 36,
|
|
459
|
+
"rolling_origin_stride": 1,
|
|
460
|
+
"prediction_format": "dataframe",
|
|
461
|
+
|
|
462
|
+
# Prediction
|
|
463
|
+
"likelihood": None,
|
|
464
|
+
"num_samples": 1,
|
|
465
|
+
"mc_dropout": False,
|
|
466
|
+
"n_jobs": -1,
|
|
467
|
+
}
|
|
468
|
+
```
|
|
469
|
+
|
|
470
|
+
---
|
|
471
|
+
|
|
472
|
+
## ๐ References
|
|
473
|
+
|
|
474
|
+
- **Darts**: [unit8co/darts](https://github.com/unit8co/darts) โ Time series forecasting library
|
|
475
|
+
- **VIEWS**: [viewsforecasting.org](https://viewsforecasting.org/) โ Violence Early-Warning System
|
|
476
|
+
|