import numpy as np import torch from torch import nn import plotly import plotly.graph_objs as go from utils.inference import apply_sigmoid_if_necessary, scale_features from utils.model_utils import * from utils.dataset_creation import inverse_scale_feature def get_result_plotting_function( inputs: np.ndarray | torch.Tensor, predictions: np.ndarray | torch.Tensor, actuals: np.ndarray | torch.Tensor, output_feature_names: list, output_colors: list, inputs_to_plot: list, model_configuration: dict, skip_scaling: list = None) -> tuple: fig_widget = go.FigureWidget() processed_predictions = apply_sigmoid_if_necessary(predictions, model_configuration) if skip_scaling is not None: features_to_scale = [feature if feature not in skip_scaling else None for feature in output_feature_names] else: features_to_scale = output_feature_names scaled_predictions = scale_features(processed_predictions, features_to_scale, model_configuration, subset_name="test") scaled_actuals = scale_features(actuals, features_to_scale, model_configuration, subset_name="test") target_features = model_configuration["feature_config"]["target_features"] ignored_features = model_configuration["feature_config"]["ignored_features"] used_targets = [feature for feature in target_features if feature['name'] not in ignored_features] output_feature_indices = list() for i, feature in enumerate(used_targets): if feature["name"] in output_feature_names: output_feature_indices.append(i) def plot_prediction_windows(index: int): window_features = inputs[index] indices = np.arange(window_features.shape[0]) actual = scaled_actuals[index] predicted = scaled_predictions[index] # clear previous traces fig_widget.data = [] fig_widget.layout.shapes = [] downsampling_rate = 1 step_size = model_configuration["preprocessing"]["window_shift"] * downsampling_rate actuals_raw_until_now = scaled_actuals[:index + 1, output_feature_indices] predicted_raw_until_now = scaled_predictions[:index + 1, output_feature_indices] out_indices = np.arange(0, len(actuals_raw_until_now) * step_size, step_size) offset = model_configuration["input_window_length"] - (index * step_size) # if find_ovs: # # calculate the number of cycles until now, only use first occurence of 0, not consecutive zeros # actual_ov_indices_raw = [x.item() for x in torch.where(actuals_raw_until_now == 0)[0].numpy()] # actual_ov_indices = [actual_ov_indices_raw[i] for i in range(len(actual_ov_indices_raw)) if # i == 0 or actual_ov_indices_raw[i] - 1 not in actual_ov_indices_raw] # num_cycles_until_now = len(actual_ov_indices) # # # find cycle starts, cycle starts are where the actuals jump from positive to negative # start_offset = 5 # cycle_start_indices = np.where(np.diff(actuals_raw_until_now) < 0)[0] + start_offset # # current_index = 0 # predicted_ov_indices = list() # for i in range(num_cycles_until_now): # first_post_0_predicted = np.where(predicted_raw_until_now[current_index:] >= 0)[0] # first_post_0_predicted = first_post_0_predicted[0].item() if len(first_post_0_predicted) > 0 else None # if first_post_0_predicted is not None: # predicted_ov_indices.append(first_post_0_predicted + current_index) # if i < len(cycle_start_indices): # current_index = cycle_start_indices[i] # # # add step sizes # actual_ov_indices = [x * step_size for x in actual_ov_indices] # predicted_ov_indices = [x * step_size for x in predicted_ov_indices] # # # plot actual and predicted ovs # for ov_index in actual_ov_indices: # fig_widget.add_vline( # x=ov_index + offset, # line=dict(color='blue', width=2, dash='dot'), # name="Actual Ovulation", # ) # # print(len(predicted_ov_indices)) # for ov_index in predicted_ov_indices: # fig_widget.add_vline( # x=ov_index + offset, # line=dict(color='red', width=2, dash='dot'), # name="Predicted Ovulation", # ) for feature in inputs_to_plot: feature_index = feature["index"] feature_name = feature["name"] fig_widget.add_scatter( x=indices, y=window_features[:, feature_index], mode="lines", name=feature_name, ) length_limiter = model_configuration["input_window_length"] // step_size print(length_limiter) for output_feature_index in output_feature_indices: color = output_colors[output_feature_index] fig_widget.add_scatter( x=(out_indices + offset)[-length_limiter:], y=actuals_raw_until_now[:, output_feature_index][-length_limiter:], mode="lines", name="Actuals Raw", line=dict(color=color, width=2, dash='dot'), ) fig_widget.add_scatter( x=(out_indices + offset)[-length_limiter:], y=predicted_raw_until_now[:, output_feature_index][-length_limiter:], mode="lines", name="Predicted Raw", line=dict(color=color, width=2), ) # add actual and predicted values print(f"actual: {actual}, predicted: {predicted}") return fig_widget, plot_prediction_windows