{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":20604,"databundleVersionId":1357052,"sourceType":"competition"}],"dockerImageVersionId":30646,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"code","source":"# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\n\n       \nimport numpy as np\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.preprocessing import MinMaxScaler\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport math\nfrom torch.nn import TransformerEncoder, TransformerEncoderLayer\n\nfrom sklearn.metrics import mean_absolute_error, mean_squared_error, r2_score\nfrom sklearn.model_selection import train_test_split","metadata":{"execution":{"iopub.status.busy":"2024-03-09T21:36:39.004340Z","iopub.execute_input":"2024-03-09T21:36:39.004795Z","iopub.status.idle":"2024-03-09T21:36:39.013751Z","shell.execute_reply.started":"2024-03-09T21:36:39.004760Z","shell.execute_reply":"2024-03-09T21:36:39.012466Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Load Data\ndf_train = pd.read_csv('../input/osic-pulmonary-fibrosis-progression/train.csv')\ndf_test = pd.read_csv('../input/osic-pulmonary-fibrosis-progression/test.csv')","metadata":{"execution":{"iopub.status.busy":"2024-03-09T21:36:39.015935Z","iopub.execute_input":"2024-03-09T21:36:39.016812Z","iopub.status.idle":"2024-03-09T21:36:39.036111Z","shell.execute_reply.started":"2024-03-09T21:36:39.016754Z","shell.execute_reply":"2024-03-09T21:36:39.034998Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train.shape, df_train.Patient.unique().shape #there are 1549 rows, but only 176 unique patients ID's. ","metadata":{"execution":{"iopub.status.busy":"2024-03-09T21:36:39.037581Z","iopub.execute_input":"2024-03-09T21:36:39.037960Z","iopub.status.idle":"2024-03-09T21:36:39.044372Z","shell.execute_reply.started":"2024-03-09T21:36:39.037932Z","shell.execute_reply":"2024-03-09T21:36:39.043566Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train.groupby('Patient').size() #the number of entries (FVC measurements) for each unique patient \n","metadata":{"execution":{"iopub.status.busy":"2024-03-09T21:36:39.046613Z","iopub.execute_input":"2024-03-09T21:36:39.047142Z","iopub.status.idle":"2024-03-09T21:36:39.060009Z","shell.execute_reply.started":"2024-03-09T21:36:39.047114Z","shell.execute_reply":"2024-03-09T21:36:39.059089Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train","metadata":{"execution":{"iopub.status.busy":"2024-03-09T21:36:39.061341Z","iopub.execute_input":"2024-03-09T21:36:39.062057Z","iopub.status.idle":"2024-03-09T21:36:39.079320Z","shell.execute_reply.started":"2024-03-09T21:36:39.062027Z","shell.execute_reply":"2024-03-09T21:36:39.078586Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"datapoints_per_patient = df_train.groupby('Patient').size()\n\n# Filter to get patients with less than 7 datapoints.\npatients_with_less_than_7 = datapoints_per_patient[datapoints_per_patient < 7]\n\n# Count patients with less than 7.\nnumber_of_patients = patients_with_less_than_7.count()\n\nprint(f'Number of patients with less than 7 datapoints: {number_of_patients}')","metadata":{"execution":{"iopub.status.busy":"2024-03-09T21:36:39.080404Z","iopub.execute_input":"2024-03-09T21:36:39.080732Z","iopub.status.idle":"2024-03-09T21:36:39.089415Z","shell.execute_reply.started":"2024-03-09T21:36:39.080705Z","shell.execute_reply":"2024-03-09T21:36:39.087812Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train = df_train.groupby('Patient').filter(lambda x: len(x) >= 7)\ndf_train.shape\n","metadata":{"execution":{"iopub.status.busy":"2024-03-09T21:36:39.091136Z","iopub.execute_input":"2024-03-09T21:36:39.091696Z","iopub.status.idle":"2024-03-09T21:36:39.111529Z","shell.execute_reply.started":"2024-03-09T21:36:39.091658Z","shell.execute_reply":"2024-03-09T21:36:39.110794Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"datapoints_per_patient = df_train.groupby('Patient').size()\n\n# Filter to get patients with less than 7 datapoints.\npatients_with_less_than_7 = datapoints_per_patient[datapoints_per_patient < 7]\n\n# Count.\nnumber_of_patients = patients_with_less_than_7.count()\n\nprint(f'Number of patients with less than 7 datapoints: {number_of_patients}')","metadata":{"execution":{"iopub.status.busy":"2024-03-09T21:36:39.112643Z","iopub.execute_input":"2024-03-09T21:36:39.113282Z","iopub.status.idle":"2024-03-09T21:36:39.122659Z","shell.execute_reply.started":"2024-03-09T21:36:39.113253Z","shell.execute_reply":"2024-03-09T21:36:39.121439Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def Seven(group):\n    if len(group) > 7:\n        return group.iloc[-7:]  # Selects the last 7 rows, all columns for each unique Patient that has > 7 entries\n    return group","metadata":{"execution":{"iopub.status.busy":"2024-03-09T21:36:39.123921Z","iopub.execute_input":"2024-03-09T21:36:39.124335Z","iopub.status.idle":"2024-03-09T21:36:39.133398Z","shell.execute_reply.started":"2024-03-09T21:36:39.124298Z","shell.execute_reply":"2024-03-09T21:36:39.132475Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train =  df_train.groupby('Patient').apply(Seven).reset_index(drop=True)","metadata":{"execution":{"iopub.status.busy":"2024-03-09T21:36:39.144185Z","iopub.execute_input":"2024-03-09T21:36:39.144536Z","iopub.status.idle":"2024-03-09T21:36:39.189815Z","shell.execute_reply.started":"2024-03-09T21:36:39.144508Z","shell.execute_reply":"2024-03-09T21:36:39.188673Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train.head(8)","metadata":{"execution":{"iopub.status.busy":"2024-03-09T21:36:39.192368Z","iopub.execute_input":"2024-03-09T21:36:39.192809Z","iopub.status.idle":"2024-03-09T21:36:39.205907Z","shell.execute_reply.started":"2024-03-09T21:36:39.192753Z","shell.execute_reply":"2024-03-09T21:36:39.204823Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train.shape","metadata":{"execution":{"iopub.status.busy":"2024-03-09T21:36:39.207459Z","iopub.execute_input":"2024-03-09T21:36:39.207968Z","iopub.status.idle":"2024-03-09T21:36:39.215404Z","shell.execute_reply.started":"2024-03-09T21:36:39.207930Z","shell.execute_reply":"2024-03-09T21:36:39.214381Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"grouped = df_train.groupby('Patient').size()\n\n# Check if all patients have 7 datapoints\nall_have_seven = all(grouped == 7)\n\nprint(f'Do all patients have exactly 7 datapoints? {all_have_seven}')\n","metadata":{"execution":{"iopub.status.busy":"2024-03-09T21:36:39.217042Z","iopub.execute_input":"2024-03-09T21:36:39.217419Z","iopub.status.idle":"2024-03-09T21:36:39.226503Z","shell.execute_reply.started":"2024-03-09T21:36:39.217384Z","shell.execute_reply":"2024-03-09T21:36:39.225480Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def process_patient_group(df_group):\n    df_group = df_group.sort_values(by='Weeks') #sort by date of FVC measurement, then get the features for each week in order\n    patient_data = {\n        'Patient': df_group['Patient'].iloc[0],\n        'Weeks': ','.join(df_group['Weeks'].astype(str)),\n        'Percent': ','.join(df_group['Percent'].astype(str)),\n        'FVC': ','.join(df_group['FVC'].astype(str)),\n        'Age': df_group['Age'].iloc[0], \n        'Sex': df_group['Sex'].iloc[0],  \n        'SmokingStatus': df_group['SmokingStatus'].iloc[0]  # Assuming smoking status is constant for all rows of the same patient\n    }\n    return pd.Series(patient_data)\n","metadata":{"execution":{"iopub.status.busy":"2024-03-09T21:36:39.229124Z","iopub.execute_input":"2024-03-09T21:36:39.229415Z","iopub.status.idle":"2024-03-09T21:36:39.237916Z","shell.execute_reply.started":"2024-03-09T21:36:39.229391Z","shell.execute_reply":"2024-03-09T21:36:39.236898Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"processed_df = df_train.groupby('Patient').apply(process_patient_group).reset_index(drop=True)\n","metadata":{"execution":{"iopub.status.busy":"2024-03-09T21:36:39.239457Z","iopub.execute_input":"2024-03-09T21:36:39.239879Z","iopub.status.idle":"2024-03-09T21:36:39.442279Z","shell.execute_reply.started":"2024-03-09T21:36:39.239851Z","shell.execute_reply":"2024-03-09T21:36:39.440601Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Display\nprocessed_df","metadata":{"execution":{"iopub.status.busy":"2024-03-09T21:36:39.444005Z","iopub.execute_input":"2024-03-09T21:36:39.444340Z","iopub.status.idle":"2024-03-09T21:36:39.459961Z","shell.execute_reply.started":"2024-03-09T21:36:39.444313Z","shell.execute_reply":"2024-03-09T21:36:39.458946Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"0.25*(7.7206 + 7.7206 + 41.0062 + 58.4860)","metadata":{"execution":{"iopub.status.busy":"2024-03-09T21:36:39.461836Z","iopub.execute_input":"2024-03-09T21:36:39.462297Z","iopub.status.idle":"2024-03-09T21:36:39.469524Z","shell.execute_reply.started":"2024-03-09T21:36:39.462258Z","shell.execute_reply":"2024-03-09T21:36:39.468402Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import Dataset\n\nclass FVCDataset(Dataset):\n    def __init__(self, dataframe):\n        self.dataframe = dataframe\n\n        # Convert FVC strings to arrays of floats\n        def convert_to_floats(s):\n            try:\n                return np.fromstring(s, dtype=float, sep=',')\n            except ValueError:\n                return np.array([])  # Handle strings that cannot be converted to float\n            \n        # Use this data just to train the scaler\n        all_fvcs = np.concatenate(dataframe['FVC'].apply(convert_to_floats)).reshape(-1, 1)\n\n        # Initialize scaler and fit it on the data\n        self.fvc_scaler = MinMaxScaler().fit(all_fvcs)\n\n    def __len__(self):\n        return len(self.dataframe)\n    \n    def __getitem__(self, idx):\n        row = self.dataframe.iloc[idx]\n        fvc = np.fromstring(row['FVC'], dtype=float, sep=',').reshape(-1, 1)\n\n        # Scale the FVC data\n        fvc_normalized = self.fvc_scaler.transform(fvc)\n\n        # Only use the first 5 FVC values as inputs\n        inputs = torch.tensor(fvc_normalized[:5], dtype=torch.float32).unsqueeze(-1)  # Add an extra dimension\n        # Use the last 2 FVC values as targets\n        targets = torch.tensor(fvc_normalized[-2:], dtype=torch.float32).squeeze()\n\n        return inputs, targets\n\n","metadata":{"execution":{"iopub.status.busy":"2024-03-09T21:36:39.470742Z","iopub.execute_input":"2024-03-09T21:36:39.471026Z","iopub.status.idle":"2024-03-09T21:36:39.481845Z","shell.execute_reply.started":"2024-03-09T21:36:39.471001Z","shell.execute_reply":"2024-03-09T21:36:39.480736Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(processed_df['Age'].unique())\nprint(processed_df['SmokingStatus'].unique())\nprint(processed_df['Sex'].unique())\n\n","metadata":{"execution":{"iopub.status.busy":"2024-03-09T21:36:39.483426Z","iopub.execute_input":"2024-03-09T21:36:39.483862Z","iopub.status.idle":"2024-03-09T21:36:39.497056Z","shell.execute_reply.started":"2024-03-09T21:36:39.483824Z","shell.execute_reply":"2024-03-09T21:36:39.495974Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset = FVCDataset(processed_df)\n","metadata":{"execution":{"iopub.status.busy":"2024-03-09T21:36:39.498766Z","iopub.execute_input":"2024-03-09T21:36:39.499181Z","iopub.status.idle":"2024-03-09T21:36:39.508682Z","shell.execute_reply.started":"2024-03-09T21:36:39.499146Z","shell.execute_reply":"2024-03-09T21:36:39.507782Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"inp, out = dataset[0]","metadata":{"execution":{"iopub.status.busy":"2024-03-09T21:36:39.509949Z","iopub.execute_input":"2024-03-09T21:36:39.510335Z","iopub.status.idle":"2024-03-09T21:36:39.516646Z","shell.execute_reply.started":"2024-03-09T21:36:39.510310Z","shell.execute_reply":"2024-03-09T21:36:39.515593Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"inp","metadata":{"execution":{"iopub.status.busy":"2024-03-09T21:36:39.518445Z","iopub.execute_input":"2024-03-09T21:36:39.518867Z","iopub.status.idle":"2024-03-09T21:36:39.527380Z","shell.execute_reply.started":"2024-03-09T21:36:39.518832Z","shell.execute_reply":"2024-03-09T21:36:39.526630Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Generate indices: an array of integers from 0 to the length of the dataset - 1\ndataset_indices = list(range(len(dataset)))\n\n# Split indices into training and validation sets\ntrain_indices, val_indices, _, _ = train_test_split(dataset_indices, dataset_indices, test_size=0.2, random_state=42)\n\n# Create subsets for training and validation using Subset\nfrom torch.utils.data import Subset\n\ntrain_dataset = Subset(dataset, train_indices)\nval_dataset = Subset(dataset, val_indices)\n\nfrom torch.utils.data import DataLoader\n\ntrain_loader = DataLoader(train_dataset, batch_size=10, shuffle=True)\nval_loader = DataLoader(val_dataset, batch_size=10, shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2024-03-09T21:36:39.531924Z","iopub.execute_input":"2024-03-09T21:36:39.532222Z","iopub.status.idle":"2024-03-09T21:36:39.539802Z","shell.execute_reply.started":"2024-03-09T21:36:39.532196Z","shell.execute_reply":"2024-03-09T21:36:39.539010Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Check first item in dataset\n#input everything up till final 3; pad dataset so that each one has the same size\n#masking prevents model from back propogating\ninputs, targets = dataset[0]\nprint(f\"Inputs: {inputs}\")\nprint(f\"Targets: {targets}\")","metadata":{"execution":{"iopub.status.busy":"2024-03-09T21:36:39.540962Z","iopub.execute_input":"2024-03-09T21:36:39.541242Z","iopub.status.idle":"2024-03-09T21:36:39.556721Z","shell.execute_reply.started":"2024-03-09T21:36:39.541217Z","shell.execute_reply":"2024-03-09T21:36:39.555908Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"input_lst = []\ntarget_lst = []\n\ndef grab_data_loop(dataset, indices):\n    input_lst = []\n    target_lst = []\n    \n    for idx in indices: \n        a, b = dataset[idx]\n        input_lst.append(a.cpu().numpy())\n        target_lst.append(b.cpu().numpy())\n        \n    X = np.array(input_lst)\n    y = np.array(target_lst)\n    \n    return X, y\n\nX_train_og, Y_train_og = grab_data_loop(dataset, train_indices)\nX_val_og, Y_val_og = grab_data_loop(dataset, val_indices)","metadata":{"execution":{"iopub.status.busy":"2024-03-09T21:36:39.558675Z","iopub.execute_input":"2024-03-09T21:36:39.559088Z","iopub.status.idle":"2024-03-09T21:36:39.623829Z","shell.execute_reply.started":"2024-03-09T21:36:39.559050Z","shell.execute_reply":"2024-03-09T21:36:39.622861Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def grab_data_loop_all(dataset, indices):\n    input_lst = []\n    target_lst = []\n    \n    for idx in indices: \n        a,b = dataset[idx]\n      \n        \n        input_lst.append(a.cpu().numpy())\n        target_lst.append(b.cpu().numpy())\n    X = np.array(input_lst)\n    y = np.array(target_lst)\n    \n    return X,y","metadata":{"execution":{"iopub.status.busy":"2024-03-09T21:36:39.625115Z","iopub.execute_input":"2024-03-09T21:36:39.625633Z","iopub.status.idle":"2024-03-09T21:36:39.631077Z","shell.execute_reply.started":"2024-03-09T21:36:39.625604Z","shell.execute_reply":"2024-03-09T21:36:39.630310Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#X_train_og is shape n_samples x week idx x n_features\n\nlast_fvc_values = X_train_og[:, -1]\n","metadata":{"execution":{"iopub.status.busy":"2024-03-09T21:36:39.632149Z","iopub.execute_input":"2024-03-09T21:36:39.632631Z","iopub.status.idle":"2024-03-09T21:36:39.641145Z","shell.execute_reply.started":"2024-03-09T21:36:39.632600Z","shell.execute_reply":"2024-03-09T21:36:39.640129Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X_train_og.shape","metadata":{"execution":{"iopub.status.busy":"2024-03-09T21:36:39.642682Z","iopub.execute_input":"2024-03-09T21:36:39.643126Z","iopub.status.idle":"2024-03-09T21:36:39.653961Z","shell.execute_reply.started":"2024-03-09T21:36:39.643088Z","shell.execute_reply":"2024-03-09T21:36:39.652928Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X_train_og, Y_train_og = grab_data_loop_all(dataset, train_indices)\nX_val_og, Y_val_og = grab_data_loop_all(dataset, val_indices)\n\nX_train_og = X_train_og.reshape(X_train_og.shape[0],-1)\nX_val_og = X_val_og.reshape(X_val_og.shape[0],-1)\nX_train_og.shape, X_val_og.shape, \n","metadata":{"execution":{"iopub.status.busy":"2024-03-09T21:36:39.655228Z","iopub.execute_input":"2024-03-09T21:36:39.655591Z","iopub.status.idle":"2024-03-09T21:36:39.723513Z","shell.execute_reply.started":"2024-03-09T21:36:39.655545Z","shell.execute_reply":"2024-03-09T21:36:39.722213Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_dictionary = {}\n\nfor i in range(X_train_og.shape[1]):  # Iterate only up to the number of columns in X_train_og\n    data_dictionary[f\"FVC_{i+1}\"] = X_train_og[:, i]  # Name the columns as FVC_1, FVC_2, etc.\n\ndataframe = pd.DataFrame(data_dictionary)\n\nimport seaborn as sns\n\ncorr_df = np.abs(dataframe.corr())\n\nsns.heatmap(corr_df)\n","metadata":{"execution":{"iopub.status.busy":"2024-03-09T21:36:39.726473Z","iopub.execute_input":"2024-03-09T21:36:39.727533Z","iopub.status.idle":"2024-03-09T21:36:40.038761Z","shell.execute_reply.started":"2024-03-09T21:36:39.727495Z","shell.execute_reply":"2024-03-09T21:36:40.037407Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataframe.shape","metadata":{"execution":{"iopub.status.busy":"2024-03-09T21:36:40.041456Z","iopub.execute_input":"2024-03-09T21:36:40.041814Z","iopub.status.idle":"2024-03-09T21:36:40.047746Z","shell.execute_reply.started":"2024-03-09T21:36:40.041785Z","shell.execute_reply":"2024-03-09T21:36:40.046711Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for j in range(dataframe.shape[0]):\n    x = dataframe.iloc[j, :]\n    y = Y_train_og[j]\n","metadata":{"execution":{"iopub.status.busy":"2024-03-09T21:36:40.049704Z","iopub.execute_input":"2024-03-09T21:36:40.050538Z","iopub.status.idle":"2024-03-09T21:36:40.067523Z","shell.execute_reply.started":"2024-03-09T21:36:40.050497Z","shell.execute_reply":"2024-03-09T21:36:40.066467Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.neural_network import MLPRegressor\n\nregr = MLPRegressor(random_state=1, max_iter=5000).fit(X_train_og, Y_train_og)\n#regr.predict(X_val_og[:2])\nregr.score(X_val_og, Y_val_og)","metadata":{"execution":{"iopub.status.busy":"2024-03-09T21:36:40.069107Z","iopub.execute_input":"2024-03-09T21:36:40.069830Z","iopub.status.idle":"2024-03-09T21:36:40.103729Z","shell.execute_reply.started":"2024-03-09T21:36:40.069791Z","shell.execute_reply":"2024-03-09T21:36:40.102612Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset.fvc_scaler.inverse_transform(Y_train_og)[:3]","metadata":{"execution":{"iopub.status.busy":"2024-03-09T21:36:40.105125Z","iopub.execute_input":"2024-03-09T21:36:40.105462Z","iopub.status.idle":"2024-03-09T21:36:40.113273Z","shell.execute_reply.started":"2024-03-09T21:36:40.105434Z","shell.execute_reply":"2024-03-09T21:36:40.112117Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.linear_model import LinearRegression, Ridge\nfrom sklearn.preprocessing import PolynomialFeatures\n\ndef train_data(a,d):\n    \n    rm = Ridge(alpha=a)\n    \n    \n        \n    pf = PolynomialFeatures(degree=d)\n\n    X_train = pf.fit_transform(X_train_og)\n    X_val = pf.transform(X_val_og)\n        \n    rm.fit(X_train, Y_train_og)\n\n    y_h_tr = rm.predict(X_train)\n    y_h_te = rm.predict(X_val)\n\n    def compute_metrics(y,y_h,scaler,datatype):\n        #Need to inverse scale to get back to normal \"dimensions\", MAE/RMSE affected, R2 the same \n        y_sc = scaler.inverse_transform(y) \n        y_h_sc = scaler.inverse_transform(y_h)\n        mae = mean_absolute_error(y_sc, y_h_sc)\n        rmse = mean_squared_error(y_sc, y_h_sc, squared=False) \n        r2 = r2_score(y_sc, y_h_sc) \n        \n        print(f\"Summary Statistics for {datatype}\\nMAE: {mae}, RMSE: {rmse}, R^2: {r2}\")\n        \n        return mae, rmse, r2\n\n    mae_tr, rmse_tr, r2_tr = compute_metrics(Y_train_og, y_h_tr, dataset.fvc_scaler, \"Training Data\")\n\n    mae_te, rmse_te, r2_te = compute_metrics(Y_val_og, y_h_te, dataset.fvc_scaler, \"Validation Data\")\n    \n    return rmse_te","metadata":{"execution":{"iopub.status.busy":"2024-03-09T21:36:40.114730Z","iopub.execute_input":"2024-03-09T21:36:40.115053Z","iopub.status.idle":"2024-03-09T21:36:40.125883Z","shell.execute_reply.started":"2024-03-09T21:36:40.115026Z","shell.execute_reply":"2024-03-09T21:36:40.124588Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.preprocessing import PolynomialFeatures\nfrom sklearn.linear_model import LinearRegression, Ridge\n\n\ndef hp_search(): \n    \n    alphas = [0, 0.0001,0.001,0.01,0.1,1]\n    degrees = [1,2,3,4,5,6,7,8]\n    #degrees = [1,2,3,4]\n    best_val_rmse = 5000000000\n    \n    for a in alphas: \n        for d in degrees: \n            print(50*\"=\")\n            print(f\"Alpha: {a}\\nDegree: {d}\")\n            rmse_te = train_data(a, d)\n            \n            if rmse_te < best_val_rmse: \n                best_val_rmse = rmse_te\n                best_a = a\n                best_d = d \n    print(f\"Optimal value of Alpha and Degree by HP Search:\\n L2Reg: {best_a} | Degree: {best_d}\")\n    return best_a, best_d\n","metadata":{"execution":{"iopub.status.busy":"2024-03-09T21:36:40.127117Z","iopub.execute_input":"2024-03-09T21:36:40.127408Z","iopub.status.idle":"2024-03-09T21:36:40.137963Z","shell.execute_reply.started":"2024-03-09T21:36:40.127383Z","shell.execute_reply":"2024-03-09T21:36:40.136913Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"a,d = hp_search()\n","metadata":{"execution":{"iopub.status.busy":"2024-03-09T21:36:40.139644Z","iopub.execute_input":"2024-03-09T21:36:40.139946Z","iopub.status.idle":"2024-03-09T21:36:40.801443Z","shell.execute_reply.started":"2024-03-09T21:36:40.139920Z","shell.execute_reply":"2024-03-09T21:36:40.799947Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data(a,d)","metadata":{"execution":{"iopub.status.busy":"2024-03-09T21:36:40.803422Z","iopub.execute_input":"2024-03-09T21:36:40.803861Z","iopub.status.idle":"2024-03-09T21:36:40.837810Z","shell.execute_reply.started":"2024-03-09T21:36:40.803823Z","shell.execute_reply":"2024-03-09T21:36:40.836227Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"","metadata":{}}]}