mirror of
https://github.com/saymrwulf/prophet.git
synced 2026-07-28 20:11:27 +00:00
Allow constant extra regressors
This commit is contained in:
parent
9a5338cbe4
commit
107f74f0f2
4 changed files with 12 additions and 11 deletions
|
|
@ -408,11 +408,11 @@ initialize_scales_fn <- function(m, initialize_scales, df) {
|
|||
m$start <- min(df$ds)
|
||||
m$t.scale <- time_diff(max(df$ds), m$start, "secs")
|
||||
for (name in names(m$extra_regressors)) {
|
||||
standardize <- m$extra_regressors[[name]]$standardize
|
||||
n.vals <- length(unique(df[[name]]))
|
||||
if (n.vals < 2) {
|
||||
stop('Regressor ', name, ' is constant.')
|
||||
standardize <- FALSE
|
||||
}
|
||||
standardize <- m$extra_regressors[[name]]$standardize
|
||||
if (standardize == 'auto') {
|
||||
if (n.vals == 2 && all(sort(unique(df[[name]])) == c(0, 1))) {
|
||||
# Don't standardize binary variables
|
||||
|
|
|
|||
|
|
@ -573,11 +573,12 @@ test_that("added_regressors", {
|
|||
fcst$yhat[1],
|
||||
fcst$trend[1] * (1 + fcst$multiplicative_terms[1]) + fcst$additive_terms[1]
|
||||
)
|
||||
# Check fails if constant extra regressor
|
||||
df$constant_feature <- 5
|
||||
# Check works with constant extra regressor of 0
|
||||
df$constant_feature <- 0
|
||||
m <- prophet()
|
||||
m <- add_regressor(m, 'constant_feature')
|
||||
expect_error(fit.prophet(m, df))
|
||||
m <- add_regressor(m, 'constant_feature', standardize = TRUE)
|
||||
m <- fit.prophet(m, df)
|
||||
expect_equal(m$extra_regressors$constant_feature$std, 1)
|
||||
})
|
||||
|
||||
test_that("set_seasonality_mode", {
|
||||
|
|
|
|||
|
|
@ -305,7 +305,7 @@ class Prophet(object):
|
|||
standardize = props['standardize']
|
||||
n_vals = len(df[name].unique())
|
||||
if n_vals < 2:
|
||||
raise ValueError('Regressor {} is constant.'.format(name))
|
||||
standardize = False
|
||||
if standardize == 'auto':
|
||||
if set(df[name].unique()) == set([1, 0]):
|
||||
# Don't standardize binary variables.
|
||||
|
|
|
|||
|
|
@ -651,12 +651,12 @@ class TestProphet(TestCase):
|
|||
fcst['trend'][0] * (1 + fcst['multiplicative_terms'][0])
|
||||
+ fcst['additive_terms'][0],
|
||||
)
|
||||
# Check fails if constant extra regressor
|
||||
df['constant_feature'] = 5
|
||||
# Check works if constant extra regressor at 0
|
||||
df['constant_feature'] = 0
|
||||
m = Prophet()
|
||||
m.add_regressor('constant_feature')
|
||||
with self.assertRaises(ValueError):
|
||||
m.fit(df.copy())
|
||||
m.fit(df)
|
||||
self.assertEqual(m.extra_regressors['constant_feature']['std'], 1)
|
||||
|
||||
def test_set_seasonality_mode(self):
|
||||
# Setting attribute
|
||||
|
|
|
|||
Loading…
Reference in a new issue