168 lines
4.9 KiB
Python
168 lines
4.9 KiB
Python
import os
|
|
from functools import partial
|
|
import math
|
|
|
|
import torch
|
|
from torch import nn
|
|
import numpy as np
|
|
|
|
from sklearn.preprocessing import StandardScaler, MinMaxScaler, RobustScaler
|
|
from torch.optim.lr_scheduler import OneCycleLR
|
|
from torch.optim import AdamW
|
|
|
|
from data_analysis.models.ov_detection.config import model_config
|
|
from utils.feature_functions import *
|
|
from utils.loss_functions import weighted_bce_loss_fn
|
|
from utils.training_utils import collate
|
|
|
|
from vsm_datascience_common import constants
|
|
|
|
feature_config = {
|
|
"feature_set_name": "full_feature_set",
|
|
"filter_criteria": {
|
|
"measurements.length": {
|
|
"$gt": constants.MEASUREMENTS_PER_DAY * 10,
|
|
"$lt": constants.MEASUREMENTS_PER_DAY * 150
|
|
},
|
|
"ends_at": {
|
|
"$exists": True,
|
|
"$ne": None
|
|
},
|
|
"$and": [
|
|
{"measurements.values": {"$not": {"$elemMatch": {"$lt": 35}}}},
|
|
{"measurements.values": {"$not": {"$elemMatch": {"$gt": 43}}}}
|
|
]
|
|
},
|
|
"augmentation": {
|
|
"use_augmentation": True,
|
|
"max_lookback": 4,
|
|
},
|
|
"feature_sets": [
|
|
"static_categorical_features",
|
|
"static_continuous_features",
|
|
"known_categorical_features",
|
|
"known_continuous_features",
|
|
"observed_categorical_features",
|
|
"observed_continuous_features",
|
|
"target_features"
|
|
],
|
|
"static_categorical_features": [
|
|
],
|
|
"static_continuous_features": [
|
|
{
|
|
"name": "average_cycle_length",
|
|
"fn": get_cycle_length_stats,
|
|
"scaler": RobustScaler
|
|
},
|
|
{
|
|
"name": "average_ovulation_day",
|
|
"fn": get_average_ovulation_day,
|
|
"scaler": RobustScaler
|
|
},
|
|
{
|
|
"name": "ovulation_std",
|
|
"fn": get_ovulation_std,
|
|
"scaler": StandardScaler
|
|
},
|
|
{
|
|
"name": "biphasic_fraction",
|
|
"fn": get_biphasic_fraction,
|
|
"scaler": MinMaxScaler
|
|
},
|
|
{
|
|
"name": "num_cycles",
|
|
"fn": get_num_cycles,
|
|
"scaler": RobustScaler
|
|
},
|
|
{
|
|
"name": "temperature_averages",
|
|
"fn": get_average_temperatures,
|
|
"scaler": StandardScaler
|
|
}
|
|
],
|
|
"known_categorical_features": [
|
|
],
|
|
"known_continuous_features": [
|
|
{
|
|
"name": "hours_from_start",
|
|
"fn": partial(get_hours_from_start, shift=0),
|
|
"scaler": MinMaxScaler,
|
|
"accumulation_fn": np.max
|
|
},
|
|
{
|
|
"name": "hour_of_day",
|
|
"fn": partial(get_hour_of_day_encoded, shift=0),
|
|
"scaler": None,
|
|
"accumulation_fn": np.max
|
|
},
|
|
{
|
|
"name": "day_of_week",
|
|
"fn": partial(get_day_of_week_encoded, shift=0),
|
|
"scaler": None,
|
|
"accumulation_fn": np.max
|
|
},
|
|
{
|
|
"name": "month_of_year",
|
|
"fn": partial(get_month_of_year_encoded, shift=0),
|
|
"scaler": None,
|
|
"accumulation_fn": np.max
|
|
}
|
|
],
|
|
"observed_categorical_features": [
|
|
],
|
|
"observed_continuous_features": [
|
|
{
|
|
"name": "temperature",
|
|
"fn": partial(get_temperature, shift=0),
|
|
"scaler": StandardScaler
|
|
},
|
|
{
|
|
"name": "rolling_average_temperature",
|
|
"fn": partial(get_rolling_average_with_padding, shift=0),
|
|
"scaler": StandardScaler
|
|
},
|
|
{
|
|
"name": "rolling_window_temperature_min",
|
|
"fn": partial(get_window_fn, window_size=constants.MEASUREMENTS_PER_DAY, fn=partial(np.min, axis=1)),
|
|
"scaler": StandardScaler
|
|
},
|
|
{
|
|
"name": "rolling_window_temperature_max",
|
|
"fn": partial(get_window_fn, window_size=constants.MEASUREMENTS_PER_DAY, fn=partial(np.max, axis=1)),
|
|
"scaler": StandardScaler
|
|
},
|
|
],
|
|
"target_features": [
|
|
# {
|
|
# "name": "fertility_probability",
|
|
# "fn": partial(get_fertility_probability, shift=0),
|
|
# "scaler": MinMaxScaler,
|
|
# "accumulation_fn": np.max
|
|
# },
|
|
# {
|
|
# "name": "ov_over_probability",
|
|
# "fn": partial(get_ov_over_probability, shift=0),
|
|
# "scaler": MinMaxScaler,
|
|
# "accumulation_fn": np.max
|
|
# },
|
|
{
|
|
"name": "days_relative_to_ov",
|
|
"fn": get_days_relative_to_ov,
|
|
"scaler": RobustScaler,
|
|
"accumulation_fn": np.max
|
|
},
|
|
{
|
|
"name": "ov_day",
|
|
"fn": get_ov_day,
|
|
"scaler": RobustScaler,
|
|
"accumulation_fn": np.max
|
|
},
|
|
{
|
|
"name": "is_biphasic",
|
|
"fn": get_is_biphasic,
|
|
"scaler": MinMaxScaler,
|
|
"accumulation_fn": np.max
|
|
}
|
|
]
|
|
}
|