Skip to content

Commit 2c4215c

Browse files
committed
change k window functionality in seasonal_onset()
1 parent 80f3918 commit 2c4215c

1 file changed

Lines changed: 33 additions & 15 deletions

File tree

‎R/seasonal_onset.R‎

Lines changed: 33 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -132,14 +132,16 @@ seasonal_onset <- function(
132132
}
133133
checkmate::reportAssertions(coll)
134134

135+
# Save the time_interval from the original tsd object
136+
time_interval <- attr(tsd, "time_interval")
137+
135138
# Add the seasons to tsd if available
136139
if (!is.null(season_start)) {
137140
tsd <- tsd |> dplyr::mutate(season = epi_calendar(.data$time, start = season_start, end = season_end))
138141
} else {
139142
tsd <- tsd |> dplyr::mutate(season = "not_defined")
140143
}
141144

142-
143145
# Define observation as cases in `tsd`.
144146
tsd <- tsd |>
145147
dplyr::mutate(observation = .data$cases)
@@ -161,19 +163,20 @@ seasonal_onset <- function(
161163
}
162164

163165
# Create the combined data frame:
164-
tsd <- dplyr::bind_rows(
165-
# If a previous season exists, use its last k-1 rows
166-
if (!is.na(prev_season)) {
166+
# If a previous season exists, use its last k-1 rows
167+
# or else use the current season
168+
if (!is.na(prev_season)) {
169+
tsd <- dplyr::bind_rows(
167170
tsd |>
168171
dplyr::filter(.data$season == prev_season) |>
169-
dplyr::slice_tail(n = k - 1)
170-
} else {
171-
tibble::tibble()
172-
},
173-
# Bind all rows from the current season
174-
tsd |>
172+
dplyr::slice_tail(n = k - 1),
173+
tsd |>
174+
dplyr::filter(.data$season == current_season)
175+
)
176+
} else {
177+
tsd <- tsd |>
175178
dplyr::filter(.data$season == current_season)
176-
)
179+
}
177180
}
178181

179182
# Extract the length of the series
@@ -225,15 +228,30 @@ seasonal_onset <- function(
225228

226229
# Estimate growth rates for all possible intervals
227230
for (i in k:n) {
228-
# Index observations for this iteration
229-
obs_iter <- tsd[(i - k + 1):i, ]
231+
232+
# Ensure continuous time steps within the k window with maximum of na_fraction_allowed
233+
current_time <- tsd$time[i]
234+
235+
# Define expected time points within the k-window
236+
expected_time <- switch(
237+
time_interval,
238+
days = current_time - lubridate::days((k - 1):0),
239+
weeks = current_time - lubridate::weeks((k - 1):0),
240+
months = lubridate::`%m-%`(current_time, lubridate::period(months = (k - 1):0))
241+
)
242+
243+
# Create complete k-window
244+
# Use match instead of left_join to reduce computation time
245+
# Missing time points are represented by NA
246+
idx <- match(expected_time, tsd$time)
247+
obs_iter <- tsd[idx, ]
248+
obs_iter$time <- expected_time
230249

231250
# Evaluate NA and zero values in windows
232251
if (sum(is.na(obs_iter$observation) | obs_iter$observation == 0) > k * na_fraction_allowed) {
233252
skipped_window[i] <- TRUE
234253
# Set fields to NA since the window is skipped
235-
growth_rates <- list(estimate = c(NA, NA, NA),
236-
fit = list(converged = FALSE))
254+
growth_rates <- list(estimate = c(NA, NA, NA), fit = list(converged = FALSE))
237255
} else {
238256
# Estimate growth rates
239257
growth_rates <- fit_growth_rate(

0 commit comments

Comments
 (0)