Surrogate¶
Cheap regression surrogates fit over a ResultsTable
for predicting observables at untested configurations.
Install via the optional extra:
trade_study.fit_surrogate(results, factors, *, method='gp', seed=0, n_estimators=200, cv_folds=5, warn_below_r2=0.0)
¶
Fit a per-observable surrogate over a :class:ResultsTable.
Rows whose score column contains NaN are dropped on a
per-observable basis (so a partially-evaluated trial still
contributes to the observables it does have).
After fitting each observable's model on all its available rows,
also computes a held-out cross-validated R^2/RMSE (#114) so callers
can tell a well-fit surrogate from one that's effectively guessing
in a sparse or noisy region of the design space -- neither backend
reports this on its own (RF's cheap OOB score isn't available for
GP, so CV is used uniformly for both, at the cost of cv_folds
extra fits per observable).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
results
|
ResultsTable
|
A :class: |
required |
factors
|
list[Factor]
|
Factor definitions used to encode |
required |
method
|
str
|
|
'gp'
|
seed
|
int
|
Random seed forwarded to the backend estimators. |
0
|
n_estimators
|
int
|
Number of trees for the |
200
|
cv_folds
|
int
|
Number of cross-validation folds used to compute
|
5
|
warn_below_r2
|
float | None
|
If not |
0.0
|
Returns:
| Type | Description |
|---|---|
SurrogateModel
|
A fitted :class: |
Raises:
| Type | Description |
|---|---|
ValueError
|
If |
Source code in src/trade_study/surrogate.py
207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 | |
trade_study.SurrogateModel(method, encoder, observable_names, models, cv_r2=dict(), cv_rmse=dict())
dataclass
¶
Fitted surrogate over a :class:ResultsTable.
Use :func:fit_surrogate to construct one. Per-observable backend
estimators are stored in models; encoding is shared across them.
Attributes:
| Name | Type | Description |
|---|---|---|
method |
str
|
|
encoder |
_FactorEncoder
|
Factor encoder used at fit time. |
observable_names |
list[str]
|
Column names of the predicted observables. |
models |
list[Any]
|
One fitted scikit-learn estimator per observable. |
cv_r2 |
dict[str, float]
|
Per-observable held-out cross-validated R^2 (#114). Missing
for an observable if too few rows were available to run CV
(fewer than 2 folds); |
cv_rmse |
dict[str, float]
|
Per-observable held-out cross-validated RMSE, companion
to |
predict(config)
¶
Predict observables for a single config.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
config
|
dict[str, Any]
|
Factor-keyed config dict. |
required |
Returns:
| Type | Description |
|---|---|
dict[str, float]
|
Mapping from observable name to predicted scalar. |
Source code in src/trade_study/surrogate.py
predict_batch(configs)
¶
Predict observables for a batch of configs.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
configs
|
Sequence[dict[str, Any]]
|
Sequence of factor-keyed config dicts. |
required |
Returns:
| Type | Description |
|---|---|
dict[str, NDArray[float64]]
|
Mapping from observable name to a length- |
dict[str, NDArray[float64]]
|
array of predictions. |
Source code in src/trade_study/surrogate.py
uncertainty(config)
¶
Predictive standard deviation per observable (GP only).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
config
|
dict[str, Any]
|
Factor-keyed config dict. |
required |
Returns:
| Type | Description |
|---|---|
dict[str, float]
|
Mapping from observable name to predictive standard deviation. |
Raises:
| Type | Description |
|---|---|
NotImplementedError
|
If the backend does not expose calibrated
uncertainties (currently anything other than |