added code
This commit is contained in:
@@ -0,0 +1,173 @@
|
||||
from tqdm import tqdm
|
||||
from vsm_datascience_common.cycle_database_connection.cycle_data import *
|
||||
from vsm_datascience_common.cycles.sequences import *
|
||||
|
||||
from utils.evaluation import aggregate_errors
|
||||
from utils.cycle_utils import get_ovulation_day
|
||||
from utils.feature_functions import *
|
||||
from utils.utils import recursive_dict_update
|
||||
|
||||
|
||||
def get_cycle_fertility_curve(cycle: dict, ov_day: int) -> np.ndarray:
|
||||
timestamps = get_timestamps(cycle)
|
||||
cycle_length = len(timestamps)
|
||||
|
||||
if ov_day is None:
|
||||
ov_index = None
|
||||
else:
|
||||
ov_timestamp = cycle["starts_at"] + timedelta(days=ov_day)
|
||||
ov_index = np.searchsorted(timestamps, ov_timestamp)
|
||||
|
||||
if ov_index is None:
|
||||
fertility_curve = np.full(cycle_length, 0.0)
|
||||
else:
|
||||
fertility_curve = get_fertility_curve(
|
||||
cycle_length,
|
||||
ov_index,
|
||||
get_fertility_probability_base_curve()[0],
|
||||
0
|
||||
)
|
||||
|
||||
return fertility_curve
|
||||
|
||||
|
||||
def get_ov_over_curve(cycle: dict, ov_day: int) -> np.ndarray:
|
||||
timestamps = get_timestamps(cycle)
|
||||
cycle_length = len(timestamps)
|
||||
|
||||
if ov_day is None:
|
||||
ov_index = None
|
||||
else:
|
||||
ov_timestamp = cycle["starts_at"] + timedelta(days=ov_day)
|
||||
ov_index = np.searchsorted(timestamps, ov_timestamp)
|
||||
|
||||
ov_over_curve = np.full(cycle_length, 0.0)
|
||||
if ov_index is not None:
|
||||
ov_over_curve[min(ov_index, cycle_length - 1):] = 1.0
|
||||
|
||||
return ov_over_curve
|
||||
|
||||
|
||||
def get_last_cycle_baseline(cycle: dict) -> tuple[np.ndarray, np.ndarray]:
|
||||
try:
|
||||
last_cycle = get_previous_cycle(cycle["_id"])
|
||||
except ValueError:
|
||||
last_cycle = None
|
||||
|
||||
current_ov_day = get_ovulation_day(cycle)
|
||||
if last_cycle is None:
|
||||
last_ov_day = 18
|
||||
else:
|
||||
last_ov_day = get_ovulation_day(last_cycle)
|
||||
|
||||
current_fertility_curve = get_cycle_fertility_curve(cycle, current_ov_day).reshape(-1, 1)
|
||||
current_ov_over_curve = get_ov_over_curve(cycle, current_ov_day).reshape(-1, 1)
|
||||
current_actual = np.concatenate((current_fertility_curve, current_ov_over_curve), axis=1)
|
||||
|
||||
last_fertility_curve = get_cycle_fertility_curve(cycle, last_ov_day).reshape(-1, 1)
|
||||
last_ov_over_curve = get_ov_over_curve(cycle, last_ov_day).reshape(-1, 1)
|
||||
last_actual = np.concatenate((last_fertility_curve, last_ov_over_curve), axis=1)
|
||||
|
||||
return last_actual, current_actual
|
||||
|
||||
|
||||
def get_population_mean_baseline(cycle: dict) -> tuple[np.ndarray, np.ndarray]:
|
||||
current_ov_day = get_ovulation_day(cycle)
|
||||
last_ov_day = 18
|
||||
|
||||
current_fertility_curve = get_cycle_fertility_curve(cycle, current_ov_day).reshape(-1, 1)
|
||||
current_ov_over_curve = get_ov_over_curve(cycle, current_ov_day).reshape(-1, 1)
|
||||
current_actual = np.concatenate((current_fertility_curve, current_ov_over_curve), axis=1)
|
||||
|
||||
last_fertility_curve = get_cycle_fertility_curve(cycle, last_ov_day).reshape(-1, 1)
|
||||
last_ov_over_curve = get_ov_over_curve(cycle, last_ov_day).reshape(-1, 1)
|
||||
last_actual = np.concatenate((last_fertility_curve, last_ov_over_curve), axis=1)
|
||||
|
||||
return last_actual, current_actual
|
||||
|
||||
|
||||
def get_user_mean_baseline(cycle: dict) -> tuple[np.ndarray, np.ndarray]:
|
||||
current_ov_day = get_ovulation_day(cycle)
|
||||
|
||||
previous_cycle = get_previous_cycles(cycle["_id"], 100)
|
||||
if previous_cycle is None:
|
||||
last_ov_day = 18
|
||||
else:
|
||||
previous_ovs = list()
|
||||
for prev_cycle in previous_cycle:
|
||||
ov_day = get_ovulation_day(prev_cycle)
|
||||
if ov_day is not None:
|
||||
previous_ovs.append(ov_day)
|
||||
if len(previous_ovs) == 0:
|
||||
last_ov_day = 18
|
||||
else:
|
||||
last_ov_day = int(np.mean(previous_ovs))
|
||||
|
||||
current_fertility_curve = get_cycle_fertility_curve(cycle, current_ov_day).reshape(-1, 1)
|
||||
current_ov_over_curve = get_ov_over_curve(cycle, current_ov_day).reshape(-1, 1)
|
||||
current_actual = np.concatenate((current_fertility_curve, current_ov_over_curve), axis=1)
|
||||
|
||||
last_fertility_curve = get_cycle_fertility_curve(cycle, last_ov_day).reshape(-1, 1)
|
||||
last_ov_over_curve = get_ov_over_curve(cycle, last_ov_day).reshape(-1, 1)
|
||||
last_actual = np.concatenate((last_fertility_curve, last_ov_over_curve), axis=1)
|
||||
|
||||
return last_actual, current_actual
|
||||
|
||||
|
||||
def get_baseline_evaluation_for_cycle(cycle: dict,
|
||||
cycle_number: int,
|
||||
predictor_fn: Callable,
|
||||
eval_fns: list) -> dict:
|
||||
current_predictions, current_actuals = predictor_fn(cycle)
|
||||
errors = dict()
|
||||
for eval_fn in eval_fns:
|
||||
if eval_fn is not None:
|
||||
eval_fn_name = eval_fn["name"]
|
||||
if eval_fn_name not in errors:
|
||||
errors[eval_fn_name] = dict()
|
||||
eval_function = eval_fn["eval_fn"]
|
||||
eval_fn_indices = [eval_fn["input_index"]] if "input_index" in eval_fn else range(
|
||||
len(current_predictions[0]))
|
||||
|
||||
for output_index in eval_fn_indices:
|
||||
if any(np.isnan(current_predictions[:, output_index])):
|
||||
continue
|
||||
|
||||
input_preds = current_predictions[:, output_index]
|
||||
input_actuals = current_actuals[:, output_index]
|
||||
error = eval_function(input_preds, input_actuals)
|
||||
|
||||
if np.isnan(error):
|
||||
# skip if error is nan
|
||||
continue
|
||||
|
||||
if f"after_{cycle_number}" not in errors[eval_fn_name]:
|
||||
errors[eval_fn_name][f"after_{cycle_number}"] = dict()
|
||||
|
||||
if output_index not in errors[eval_fn_name][f"after_{cycle_number}"]:
|
||||
errors[eval_fn_name][f"after_{cycle_number}"][output_index] = list()
|
||||
|
||||
errors[eval_fn_name][f"after_{cycle_number}"][output_index].append(error)
|
||||
|
||||
return errors
|
||||
|
||||
|
||||
def get_errors_for_users(user_ids: list,
|
||||
key_stats: dict,
|
||||
predictor_fn: Callable,
|
||||
eval_fns: list,
|
||||
aggregate: bool = True,
|
||||
show_progress: bool = False) -> dict:
|
||||
errors = dict()
|
||||
for key in tqdm(user_ids, disable=not show_progress):
|
||||
user_cycle_ids = [c["cycle_id"] for c in key_stats["by_key"][key]["cycle_stats"]]
|
||||
user_cycles = [get_cycle_by_id(cycle_id) for cycle_id in user_cycle_ids]
|
||||
# sort by starts_at
|
||||
user_cycles = sorted(user_cycles, key=lambda c: c["starts_at"])
|
||||
for i, cycle in enumerate(user_cycles):
|
||||
current_errors = get_baseline_evaluation_for_cycle(cycle, i, predictor_fn, eval_fns)
|
||||
errors = recursive_dict_update(errors, current_errors)
|
||||
|
||||
if aggregate:
|
||||
errors = aggregate_errors(errors, eval_fns)
|
||||
return errors
|
||||
Reference in New Issue
Block a user